mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
refactor: pass outcome contexts instead of a record flag and reuse the shared helpers
This commit is contained in:
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user