diff --git a/routstr/core/terminal_outcomes.py b/routstr/core/terminal_outcomes.py index a1470324..ba4ad031 100644 --- a/routstr/core/terminal_outcomes.py +++ b/routstr/core/terminal_outcomes.py @@ -3,6 +3,7 @@ from __future__ import annotations import time from dataclasses import dataclass +from ..payment.usage import NormalizedUsage, normalize_usage from .logging import get_logger from .terminal_outcome_writer import ( TerminalOutcomeWriter, @@ -52,8 +53,6 @@ def record_terminal_outcome( """ try: if usage is not None: - from ..payment.usage import NormalizedUsage, normalize_usage - counted = normalize_usage(usage) or NormalizedUsage() input_tokens = counted.input_tokens output_tokens = counted.output_tokens diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index fa230f5e..efb45b63 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -4249,7 +4249,7 @@ class BaseUpstreamProvider: request_id: str | None = None, model_obj: Model | None = None, request_body: bytes | None = None, - record_outcome: bool = True, + terminal_outcome: TerminalOutcomeContext | None = None, ) -> StreamingResponse: """Handle streaming response for X-Cashu payment, calculating refund if needed. @@ -4284,9 +4284,7 @@ class BaseUpstreamProvider: usage_estimator = MissingUsageEstimator(request_body, model_obj) refund_amount_sent = 0 settlement_failed = False - outcome_state = TerminalOutcomeState( - terminal_outcome_context(request_id, model_obj) if record_outcome else None - ) + outcome_state = TerminalOutcomeState(terminal_outcome) # Stats observe both SSE prefix forms; billing keeps its existing parse. observe_terminal_sse_bytes(outcome_state, b"", content_str.encode(), final=True) @@ -4480,7 +4478,7 @@ class BaseUpstreamProvider: request_id: str | None = None, model_obj: Model | None = None, request_body: bytes | None = None, - record_outcome: bool = True, + terminal_outcome: TerminalOutcomeContext | None = None, ) -> Response: """Handle non-streaming response for X-Cashu payment, calculating refund if needed. @@ -4501,11 +4499,7 @@ class BaseUpstreamProvider: try: response_json = json.loads(content_str) - outcome_state = TerminalOutcomeState( - terminal_outcome_context(request_id, model_obj) - if record_outcome - else None - ) + outcome_state = TerminalOutcomeState(terminal_outcome) outcome_state.observe(response_json) self._apply_provider_field(response_json) _apply_estimated_usage( @@ -4669,7 +4663,7 @@ class BaseUpstreamProvider: request_id: str | None = None, model_obj: Model | None = None, request_body: bytes | None = None, - record_outcome: bool = True, + terminal_outcome: TerminalOutcomeContext | None = None, ) -> StreamingResponse | Response: """Handle chat completion response for X-Cashu payment, detecting streaming vs non-streaming. @@ -4717,7 +4711,7 @@ class BaseUpstreamProvider: request_id=request_id, model_obj=model_obj, request_body=request_body, - record_outcome=record_outcome, + terminal_outcome=terminal_outcome, ) else: return await self.handle_x_cashu_non_streaming_response( @@ -4730,7 +4724,7 @@ class BaseUpstreamProvider: request_id=request_id, model_obj=model_obj, request_body=request_body, - record_outcome=record_outcome, + terminal_outcome=terminal_outcome, ) except Exception as e: @@ -4934,7 +4928,11 @@ class BaseUpstreamProvider: request_id=getattr(request.state, "request_id", None), model_obj=model_obj, request_body=request_body, - record_outcome=not path.endswith("messages/count_tokens"), + terminal_outcome=None + if path.endswith("messages/count_tokens") + else terminal_outcome_context( + getattr(request.state, "request_id", None), model_obj + ), ) if isinstance(result, StreamingResponse) and not response.is_closed: return attach_upstream_stream_owner(result, response, client) diff --git a/routstr/upstream/ehbp.py b/routstr/upstream/ehbp.py index 6f58b570..d5351905 100644 --- a/routstr/upstream/ehbp.py +++ b/routstr/upstream/ehbp.py @@ -43,7 +43,6 @@ from ..core.exceptions import EhbpTimeoutError, UpstreamError from ..core.settings import settings from ..core.terminal_outcomes import ( TerminalOutcomeContext, - cashu_retained_msats, mark_terminal_outcome_loss, record_terminal_outcome, ) @@ -62,6 +61,10 @@ from ..wallet import ( recieve_token, send_token, ) +from .terminal_outcome_tracking import ( + record_x_cashu_terminal_outcome, + terminal_outcome_context, +) from .tinfoil_trailer import TrailerResponse, forward_with_trailer logger = get_logger(__name__) @@ -1002,10 +1005,8 @@ async def forward_ehbp_request( trailer (streaming). Usage is captured from both response headers and HTTP trailers via an h11-based client (httpx silently discards trailers). """ - terminal_outcome = TerminalOutcomeContext( - outcome_id=getattr(request.state, "request_id", None), - model_identifier=model_obj.canonical_slug or model_obj.id, - served_model_identifier=model_obj.forwarded_model_id or model_obj.id, + terminal_outcome = terminal_outcome_context( + getattr(request.state, "request_id", None), model_obj ) target = upstream.get_ehbp_forwarding_target(path, model_obj) # type: ignore[attr-defined] @@ -1281,11 +1282,7 @@ async def forward_ehbp_x_cashu_request( client because httpx silently discards them. """ request_id = getattr(request.state, "request_id", None) - terminal_outcome = TerminalOutcomeContext( - outcome_id=request_id, - model_identifier=model_obj.canonical_slug or model_obj.id, - served_model_identifier=model_obj.forwarded_model_id or model_obj.id, - ) + terminal_outcome = terminal_outcome_context(request_id, model_obj) amount = 0 unit = "msat" mint: str | None = None @@ -1465,22 +1462,13 @@ async def forward_ehbp_x_cashu_request( raise persisted_refund_amount = refund_amount - revenue_msats = cashu_retained_msats( - amount, - unit, + record_x_cashu_terminal_outcome( + terminal_outcome, + cost_info, + amount=amount, + unit=unit, refund_amount=persisted_refund_amount, ) - if revenue_msats is not None: - record_terminal_outcome( - terminal_outcome, - input_tokens=cost_info.get("input_tokens", 0), - output_tokens=cost_info.get("output_tokens", 0), - cache_read_input_tokens=cost_info.get("cache_read_input_tokens", 0), - cache_creation_input_tokens=cost_info.get( - "cache_creation_input_tokens", 0 - ), - revenue_msats=revenue_msats, - ) async def _stream_body_xcashu() -> AsyncIterator[bytes]: yield resp.body diff --git a/tests/unit/test_x_cashu_cost_sats.py b/tests/unit/test_x_cashu_cost_sats.py index 55460a0f..7b7873b8 100644 --- a/tests/unit/test_x_cashu_cost_sats.py +++ b/tests/unit/test_x_cashu_cost_sats.py @@ -137,6 +137,7 @@ async def test_non_streaming_cost_sats_value_rounds_down() -> None: unit="sat", max_cost_for_model=10000, request_id="sat-rounding-request", + terminal_outcome=TerminalOutcomeContext("sat-rounding-request", None), ) body = json.loads(response.body) @@ -286,6 +287,7 @@ async def test_streaming_no_space_error_event_is_not_recorded() -> None: unit="sat", max_cost_for_model=10_000, request_id="no-space-error", + terminal_outcome=TerminalOutcomeContext("no-space-error", None), ) record.assert_not_called() @@ -329,6 +331,7 @@ async def test_streaming_no_space_usage_is_still_recorded_as_reported() -> None: unit="sat", max_cost_for_model=10_000, request_id="no-space-recorded", + terminal_outcome=TerminalOutcomeContext("no-space-recorded", None), ) context = record.call_args.args[0] @@ -365,6 +368,7 @@ async def test_native_messages_stream_keeps_input_usage_in_stats() -> None: unit="msat", max_cost_for_model=100, request_id="messages-request", + terminal_outcome=TerminalOutcomeContext("messages-request", None), ) await _collect_streaming(response)