diff --git a/routstr/core/terminal_outcomes.py b/routstr/core/terminal_outcomes.py index ba4ad031..03c52cad 100644 --- a/routstr/core/terminal_outcomes.py +++ b/routstr/core/terminal_outcomes.py @@ -1,5 +1,8 @@ from __future__ import annotations +import hashlib +import hmac +import secrets import time from dataclasses import dataclass @@ -14,6 +17,10 @@ from .terminal_outcome_writer import ( logger = get_logger(__name__) +# Rows must not share an id with request logs or the x-routstr-request-id +# header. The key never leaves this process, so ids cannot be recomputed. +_OUTCOME_ID_KEY = secrets.token_bytes(32) + # Far above any real request, and low enough that a day's sums stay JSON-safe. _MAX_TOKENS = 2**31 - 1 _MAX_REVENUE_MSATS = 2**40 @@ -88,7 +95,7 @@ def record_terminal_outcome( return terminal_outcome_writer.submit( _QueuedOutcome( - outcome_id=context.outcome_id, + outcome_id=_outcome_id(context.outcome_id), terminal_at_ms=timestamp, terminal_day=terminal_day, model_identifier=context.model_identifier, @@ -110,6 +117,10 @@ def record_terminal_outcome( pass +def _outcome_id(request_id: str) -> str: + return hmac.new(_OUTCOME_ID_KEY, request_id.encode(), hashlib.sha256).hexdigest() + + def mark_terminal_outcome_loss(reason: str) -> None: try: terminal_outcome_writer.declare_loss(reason) diff --git a/tests/unit/test_terminal_outcomes.py b/tests/unit/test_terminal_outcomes.py index 4b6dbe10..9fc878ee 100644 --- a/tests/unit/test_terminal_outcomes.py +++ b/tests/unit/test_terminal_outcomes.py @@ -6,6 +6,7 @@ from contextlib import asynccontextmanager from dataclasses import dataclass from datetime import UTC, date, datetime, timedelta from pathlib import Path +from unittest.mock import MagicMock import pytest from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine @@ -470,7 +471,9 @@ def test_record_wrapper_never_raises_on_invalid_or_failed_submission( def submit(self, outcome: _QueuedOutcome) -> bool: self.submissions.append(outcome) - if outcome.outcome_id == "request-submit-error": + if outcome.outcome_id == outcomes_module._outcome_id( + "request-submit-error" + ): raise RuntimeError("submission failed") return True @@ -526,6 +529,27 @@ def test_record_wrapper_never_raises_on_invalid_or_failed_submission( assert writer.submissions[0].model_identifier is None +def test_outcome_rows_do_not_carry_the_request_id( + monkeypatch: pytest.MonkeyPatch, +) -> None: + writer = MagicMock() + monkeypatch.setattr(outcomes_module, "terminal_outcome_writer", writer) + for _ in range(2): + record_terminal_outcome( + TerminalOutcomeContext("request-a", "author/model"), + input_tokens=1, + output_tokens=1, + cache_read_input_tokens=0, + cache_creation_input_tokens=0, + revenue_msats=1, + ) + + first, second = (call.args[0] for call in writer.submit.call_args_list) + assert "request-a" not in first.outcome_id + # A second record of one request still collides instead of double counting. + assert first.outcome_id == second.outcome_id + + def test_cashu_retained_msats_uses_exact_persisted_units( monkeypatch: pytest.MonkeyPatch, ) -> None: