refactor: pass outcome contexts instead of a record flag and reuse the shared helpers

This commit is contained in:
Ashen
2026-10-01 11:21:53 +05:30
parent dadb9c5d74
commit c9f5dfc0b5
4 changed files with 29 additions and 40 deletions
+1 -2
View File
@@ -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
+12 -14
View File
@@ -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)
+12 -24
View File
@@ -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
+4
View File
@@ -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)