diff --git a/CHANGELOG.md b/CHANGELOG.md index 1d5a60c3f..fabff7ba1 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -12,6 +12,91 @@ and this project uses [Semantic Versioning](https://semver.org/spec/v2.0.0.html) ### Fixed +- The cost ledger no longer fabricates a $0.00 price for an unpriced + provider/model. `PriceBook.compute_cost()` now returns a + `(cost_amount, currency_code, price_known)` 3-tuple; `UsageRecord` gains a + `price_known: bool` field (persisted through a new `usage_price_knowledge` + satellite table, joined the same way `usage_measurements` already is) so a + measured-but-unpriced request is distinguishable from a genuinely free one. + `CostLedger.rollup()`/`.report()`/`.total()` add an additive + `cost_amount_by_status`/`record_count_by_status` breakdown + (measured/estimated/unavailable) alongside the existing flat `cost_amount` + total, so measured, estimated, and unavailable-priced spend are no longer + opaquely blended into one authoritative-looking number. +- `CostLedger.rollup()`/`.report()` add the same treatment one level further: + an additive `cost_amount_by_price_status`/`record_count_by_price_status` + breakdown (`known`/`unknown`) alongside the measured/estimated/unavailable + one, so an unpriced request's spend is visible even after rolling many + records up into one bucket. `cheapest_upstream()` no longer treats an + unpriced candidate as free (cost `0`) when selecting the lowest-cost + upstream — an unknown price is excluded from the comparison entirely + rather than winning it by default; `None` is returned when no candidate + has a known price. +- (Devin review on #956) `SqlLedgerStore._append_locked()`'s satellite + writes (`usage_measurements`, `usage_price_knowledge`, attribution) now + only run when the parent `llm_usage_records` insert is actually accepted. + A retried `usage_record_id` whose parent insert is correctly rejected as + a duplicate previously still ran these satellite inserts unconditionally + — for a parent row that predates one of these tables (e.g. an + upgrade-migrated row with no `usage_price_knowledge` child, intentionally + read as price-unknown) that silently backfilled its provenance from the + retry's current price/measurement state instead of what actually priced + the original spend. `append()` is now a true no-op on a rejected + duplicate. +- `price_known` now propagates from the ledger into every downstream usage + surface: `CostRoutingCoordinator.complete()`'s sync/provider-request cost + dicts and their per-currency components, `record_stream_usage()`, + `retrieve_batch()`'s per-item results, and `embeddings_batch_document()`. + An unpriced request's `cost_amount` is `null` rather than a silent `0` + wherever it surfaces, not just in the ledger's own rollups. + `PriceBook.compute_cost()` treats zero token usage as the one exception: + zero tokens cost zero regardless of whether the provider/model's price is + known (zero times any finite price is still zero), so a cache hit — whose + synthetic `("cache", "response")` provider/model never has a price row — + is always `price_known=True`, `cost_amount=0.0` rather than being wrongly + reported as an unpriced request. +- `retrieve_batch()`'s `item.prompt_tokens or None` treated a legitimately + reported zero token count the same as a missing one (Python's falsy-zero), + so a genuinely confirmed zero-usage batch item was silently downgraded to + an unmeasured estimate instead of staying a real, priced `measured` `0`. + `PgLlmBatchBackend.retrieve()` and `BatchResultItem` now carry an explicit + `usage_valid` tri-state (confirmed-typed non-negative counts vs. missing/ + malformed usage), and `retrieve_batch()` reads that instead of relying on + truthiness. That same estimation fallback also passed a hardcoded + empty-content placeholder instead of the request actually submitted, so a + large prompt whose provider marked usage invalid was undercounted to + near-zero prompt tokens — understating batch cost. (CodeRabbit review) + `submit_batch()` now computes a real prompt-token estimate per + `custom_id` before submission and stores it as a `prompt_token_estimates` + field on the durable `BatchJob` record itself, rather than the raw + submitted messages (Devin review: a batch registry can be Valkey-backed + with a multi-day retention shared across processes, and a submitted + prompt may be ZDR-flagged or otherwise sensitive — a token count carries + no reconstructable prompt content). Publishing it on the existing + `BatchJob` write also means an accepted job still has exactly one + publication write, so a metadata-only estimate can never orphan an + already-accepted (and possibly billed) backend job behind a raised + exception with no job id ever returned to the caller. `retrieve_batch()` + reads that stored estimate instead of falling back to an empty + placeholder. (Devin review) A job accepted before this fix has no + `prompt_token_estimates` at all, and its original request never lived + anywhere durable that a post-fix retrieval could still read — so + `retrieve_batch()` now also reads a legacy, pre-fix `batch_requests` + registry entry (still populated only for jobs submitted before this + change; nothing writes new entries there anymore) whenever a custom_id + has no stored estimate, computes the real prompt-token count from it the + same way a fresh submission would have, and persists that estimate back + onto the job so a re-retrieval never repeats the lookup. That legacy + lookup initially gated on whether `job.prompt_token_estimates` was + non-empty at all — wrong once any one custom_id had already picked up an + estimate (including from an earlier partial retrieval of the same job), + since every other still-unestimated custom_id would then silently stop + being looked up for the rest of that job's lifetime. It now gates on + whether the current retrieval actually has an item that still needs the + legacy lookup, so a legacy job's estimates can be filled in correctly + across as many partial retrievals as it takes; a job that has never + needed the legacy path (everything submitted after this fix) still never + touches that registry. - `CostRoutingCoordinator._record_race_endpoint_usage()` no longer silently drops a completed, billable race-loser call's spend when its usage payload can't be parsed. It now writes a `measurement_status="unavailable"` ledger diff --git a/contextual_orchestrator/batch_routing.py b/contextual_orchestrator/batch_routing.py index 82ce36600..2292c2f73 100644 --- a/contextual_orchestrator/batch_routing.py +++ b/contextual_orchestrator/batch_routing.py @@ -131,8 +131,8 @@ def cheapest_upstream( Cost-optimising upstream selection for load balancing: given candidate provider/model pairs, price each against the configurable price table for a representative request shape and return the cheapest. Unpriced candidates - cost ``0`` and are treated as free (explicit, so a missing price is visible - rather than silently expensive). Ties keep input order. + are excluded because an unknown price is not free. Ties keep input order; + ``None`` is returned when no candidate has a known price. """ if not candidates: return None @@ -141,9 +141,11 @@ def cheapest_upstream( for candidate in candidates: provider = candidate.get("provider", "") model = candidate.get("model", "") - cost, _currency = price_book.compute_cost( + cost, _currency, price_known = price_book.compute_cost( provider, model, assumed_prompt_tokens, assumed_completion_tokens ) + if not price_known: + continue if best_cost is None or cost < best_cost: best_cost = cost best = candidate @@ -190,6 +192,9 @@ class BatchJob: # HTTP callers bind this opaque digest to the authenticated principal; # library-only jobs may remain unowned for standalone use. owner_id: Optional[str] = None + # Prompt-token fallback estimates are safe metadata, stored atomically with + # the job handle rather than retaining submitted prompt text. + prompt_token_estimates: Dict[str, int] = field(default_factory=dict) @dataclass @@ -203,6 +208,7 @@ class BatchResultItem: attribution: Dict[str, Any] = field(default_factory=dict) model: str = "contextual-orchestrator" mode: str = "auto" + usage_valid: Optional[bool] = None class BatchBackend(Protocol): @@ -403,17 +409,26 @@ async def _download() -> Dict[str, Any]: body = (entry.get("response") or {}).get("body", {}) answer = _extract_answer(body) usage = body.get("usage", {}) or {} + raw_prompt_tokens = usage.get("prompt_tokens") + raw_completion_tokens = usage.get("completion_tokens") + usage_valid = ( + type(raw_prompt_tokens) is int + and raw_prompt_tokens >= 0 + and type(raw_completion_tokens) is int + and raw_completion_tokens >= 0 + ) raw_request = tracked.get(custom_id) request = BatchRequest(**raw_request) if raw_request else None items.append( BatchResultItem( custom_id=custom_id, answer=answer, - prompt_tokens=int(usage.get("prompt_tokens", 0)), - completion_tokens=int(usage.get("completion_tokens", 0)), + prompt_tokens=raw_prompt_tokens if usage_valid else 0, + completion_tokens=raw_completion_tokens if usage_valid else 0, attribution=dict(request.attribution) if request else {}, model=request.model if request else "contextual-orchestrator", mode=request.mode if request else "auto", + usage_valid=usage_valid, ) ) return items diff --git a/contextual_orchestrator/cost_ledger.py b/contextual_orchestrator/cost_ledger.py index 1622fcda1..30f2daf00 100644 --- a/contextual_orchestrator/cost_ledger.py +++ b/contextual_orchestrator/cost_ledger.py @@ -23,7 +23,8 @@ DB object names are two-or-more-word snake_case per the repository convention: ``llm_usage_records``, ``cost_attribution_dimensions``, ``llm_price_entries``, -``cost_attribution_values``, and ``usage_record_attributions``. +``cost_attribution_values``, ``usage_record_attributions``, +``usage_measurements``, and ``usage_price_knowledge``. """ from __future__ import annotations @@ -218,15 +219,26 @@ def compute_cost( model: str, prompt_tokens: int, completion_tokens: int, - ) -> tuple[float, str]: - """Return ``(cost_amount, currency_code)`` for a request. + ) -> tuple[float, str, bool]: + """Return ``(cost_amount, currency_code, price_known)`` for a request. An unpriced provider/model yields ``0.0`` in the default currency so - recording never fails on a missing price row. + recording never fails on a missing price row — but ``price_known`` is + ``False`` for that case so callers can tell "priced at zero" apart + from "we do not know the price" instead of silently fabricating a + free price for an unpriced model. + + Zero usage is the one exception: zero tokens cost zero regardless of + whether this provider/model's per-token price is known, since zero + times any finite price is still zero. A cache hit (and any other + genuinely zero-token record, priced or not) is therefore always + ``price_known=True`` -- there is no unknown quantity left to guess. """ entry = self.get_price(provider, model) + if prompt_tokens == 0 and completion_tokens == 0: + return 0.0, entry.currency_code if entry is not None else self.default_currency, True if entry is None: - return 0.0, self.default_currency + return 0.0, self.default_currency, False prompt_cost = (Decimal(prompt_tokens) / Decimal(1000)) * Decimal( str(entry.prompt_price_per_1k) ) @@ -236,13 +248,18 @@ def compute_cost( total = (prompt_cost + completion_cost).quantize( Decimal("0.000001"), rounding=ROUND_HALF_UP ) - return float(total), entry.currency_code + return float(total), entry.currency_code, True # --------------------------------------------------------------------------- # Usage records + stores # --------------------------------------------------------------------------- +# The only recognized provenance labels for a recorded usage row's token +# counts. Shared between ``record_usage``'s validation and the rollup/total +# per-status breakdown so the two never drift apart. +MEASUREMENT_STATUSES: tuple[str, ...] = ("measured", "estimated", "unavailable") + @dataclass class UsageRecord: @@ -261,6 +278,12 @@ class UsageRecord: cost_amount: float currency_code: str measurement_status: str = "measured" + # Whether ``cost_amount`` reflects a real configured price rather than the + # fallback zero used when a provider/model has no price row. Defaults to + # ``True`` so callers that set ``cost_amount`` explicitly (bypassing + # ``PriceBook.compute_cost``) are not force-changed by this field's + # addition. + price_known: bool = True attribution: AttributionDimensions = field(default_factory=AttributionDimensions) def as_dict(self) -> Dict[str, Any]: @@ -281,6 +304,7 @@ def as_dict(self) -> Dict[str, Any]: "cost_amount": self.cost_amount, "currency_code": self.currency_code, "measurement_status": self.measurement_status, + "price_known": self.price_known, } row.update( { @@ -334,6 +358,7 @@ def from_record( "contextual_orchestrator.usage_record_id": record.usage_record_id, "contextual_orchestrator.request_channel": record.request_channel, "contextual_orchestrator.usage.export_state": export_state, + "contextual_orchestrator.usage.price_known": record.price_known, "contextual_orchestrator.usage.measurement_status": record.measurement_status, } if record.workflow_run_id: @@ -352,9 +377,10 @@ def from_record( "gen_ai.usage.input_tokens": float(record.prompt_tokens), "gen_ai.usage.output_tokens": float(record.completion_tokens), "gen_ai.usage.total_tokens": float(record.total_tokens), - "gen_ai.usage.cost": float(record.cost_amount), } ) + if record.price_known and record.cost_amount is not None: + metrics["gen_ai.usage.cost"] = float(record.cost_amount) if error_type: metrics["contextual_orchestrator.usage.export_failures"] = 1.0 return cls( @@ -725,6 +751,14 @@ def __len__(self) -> int: ON DELETE CASCADE ); +CREATE TABLE IF NOT EXISTS usage_price_knowledge ( + usage_record_id TEXT PRIMARY KEY, + price_known INTEGER NOT NULL, + CONSTRAINT usage_price_knowledge_record_foreign_key + FOREIGN KEY (usage_record_id) REFERENCES llm_usage_records(usage_record_id) + ON DELETE CASCADE +); + CREATE TABLE IF NOT EXISTS cost_attribution_values ( dimension_name TEXT NOT NULL, dimension_value TEXT NOT NULL, @@ -783,6 +817,7 @@ def __len__(self) -> int: "cost_amount", "currency_code", "measurement_status", + "price_known", ) _RELATIONAL_ATTRIBUTION_COLUMNS = { "account": "account_name", @@ -874,6 +909,16 @@ def __len__(self) -> int: "VALUES (%s, %s) ON CONFLICT (usage_record_id) DO NOTHING" ), } +_USAGE_PRICE_KNOWLEDGE_INSERT_SQL = { + "qmark": ( + "INSERT INTO usage_price_knowledge (usage_record_id, price_known) " + "VALUES (?, ?) ON CONFLICT (usage_record_id) DO NOTHING" + ), + "pyformat": ( + "INSERT INTO usage_price_knowledge (usage_record_id, price_known) " + "VALUES (%s, %s) ON CONFLICT (usage_record_id) DO NOTHING" + ), +} _USAGE_SELECT_SQL = ( "SELECT u.usage_record_id, u.created_at, u.workflow_run_id, u.request_channel, " "u.route_mode, u.provider_name, u.model_name, " @@ -884,15 +929,18 @@ def __len__(self) -> int: "COALESCE(MAX(CASE WHEN a.dimension_name = 'group' THEN a.dimension_value END), 'unattributed') AS group_name, " "COALESCE(MAX(CASE WHEN a.dimension_name = 'company' THEN a.dimension_value END), 'unattributed') AS company_name, " "u.prompt_tokens, u.completion_tokens, u.total_tokens, u.cost_amount, u.currency_code, " - "COALESCE(m.measurement_status, 'unavailable') AS measurement_status " + "COALESCE(m.measurement_status, 'unavailable') AS measurement_status, " + "COALESCE(k.price_known, 0) AS price_known " "FROM llm_usage_records AS u " "LEFT JOIN usage_record_attributions AS a ON a.usage_record_id = u.usage_record_id " - "LEFT JOIN usage_measurements AS m ON m.usage_record_id = u.usage_record_id" + "LEFT JOIN usage_measurements AS m ON m.usage_record_id = u.usage_record_id " + "LEFT JOIN usage_price_knowledge AS k ON k.usage_record_id = u.usage_record_id" ) _USAGE_GROUP_ORDER_SQL = ( " GROUP BY u.usage_record_id, u.created_at, u.workflow_run_id, u.request_channel, " "u.route_mode, u.provider_name, u.model_name, u.prompt_tokens, " - "u.completion_tokens, u.total_tokens, u.cost_amount, u.currency_code, m.measurement_status " + "u.completion_tokens, u.total_tokens, u.cost_amount, u.currency_code, m.measurement_status, " + "k.price_known " "ORDER BY u.created_at, u.usage_record_id" ) _USAGE_QUERY_SQL = { @@ -1016,6 +1064,13 @@ def _copy_flattened_usage_rows(self, cur: Any, source_table: str) -> None: _USAGE_MEASUREMENT_INSERT_SQL[self._paramstyle], (row["usage_record_id"], "unavailable"), ) + # A flattened legacy row predates this signal entirely, so + # whether its stored cost reflects a real price is genuinely + # unknown — mark it unknown (0) rather than assume it was known. + cur.execute( + _USAGE_PRICE_KNOWLEDGE_INSERT_SQL[self._paramstyle], + (row["usage_record_id"], 0), + ) self._insert_normalized_attribution(cur, row) def _create_schema(self) -> None: @@ -1134,11 +1189,26 @@ def _append_locked(self, record: UsageRecord) -> bool: tuple(row.get(column) for column in _CORE_USAGE_COLUMNS), ) accepted = getattr(cur, "rowcount", 1) != 0 - cur.execute( - _USAGE_MEASUREMENT_INSERT_SQL[self._paramstyle], - (row["usage_record_id"], row.get("measurement_status", "unavailable")), - ) - self._insert_normalized_attribution(cur, row) + # A rejected duplicate parent insert (a retried usage_record_id) + # must stay fully idempotent: skipping the satellite writes below + # keeps a duplicate append a true no-op. Without this guard, a + # retry whose parent row already exists but predates one of + # these satellite tables (e.g. a pre-upgrade row with no + # usage_price_knowledge child, intentionally read as + # price-unknown) would insert that child using the retry's + # *current* price/measurement state -- silently relabeling a + # historical unknown-price row's provenance from an unrelated + # later call, not from what actually priced that original spend. + if accepted: + cur.execute( + _USAGE_MEASUREMENT_INSERT_SQL[self._paramstyle], + (row["usage_record_id"], row.get("measurement_status", "unavailable")), + ) + cur.execute( + _USAGE_PRICE_KNOWLEDGE_INSERT_SQL[self._paramstyle], + (row["usage_record_id"], 1 if row.get("price_known", True) else 0), + ) + self._insert_normalized_attribution(cur, row) if outer_transaction: cur.execute("RELEASE SAVEPOINT usage_record_append") else: @@ -1205,6 +1275,18 @@ def _within_window(created_at: int, start: Optional[int], end: Optional[int]) -> return True +def _measurement_status_of(row: Dict[str, Any]) -> str: + """Return a row's recognized measurement status, defaulting to unavailable. + + Covers a row whose ``measurement_status`` is missing entirely (a caller + or store bypassing the normal write path) the same way a genuinely + unrecognized value is covered: fail closed to the most conservative + label rather than silently attributing it to ``measured``. + """ + status = row.get("measurement_status") + return status if status in MEASUREMENT_STATUSES else "unavailable" + + # --------------------------------------------------------------------------- # The ledger # --------------------------------------------------------------------------- @@ -1290,10 +1372,10 @@ def record_usage( if provider: dims.upstream_api = provider - cost_amount, currency = self.price_book.compute_cost( + cost_amount, currency, price_known = self.price_book.compute_cost( provider, model, prompt_tokens, completion_tokens ) - if measurement_status not in {"measured", "estimated", "unavailable"}: + if measurement_status not in MEASUREMENT_STATUSES: raise ValueError("measurement_status must be measured, estimated, or unavailable") record = UsageRecord( usage_record_id=usage_record_id or f"usage_{uuid.uuid4().hex}", @@ -1309,6 +1391,7 @@ def record_usage( cost_amount=cost_amount, currency_code=currency, measurement_status=measurement_status, + price_known=price_known, attribution=dims, ) accepted = False @@ -1390,6 +1473,15 @@ def rollup( ``dimension`` is one of the attribution dimension names (``account``, ``service``, ``upstream_api``/``provider``, ``model_name``, ``team``, ``group``, ``company``). Returns a mapping of dimension value to totals. + + Each bucket's ``cost_amount`` sums every row regardless of provenance + (unchanged, backward-compatible meaning). Alongside it, + ``cost_amount_by_status``/``record_count_by_status`` break the same + bucket down by row ``measurement_status`` (``measured``, ``estimated``, + ``unavailable``) so a caller can tell how much of a total is real + provider-measured spend versus an estimate or an unpriced/unmeasured + row. The corresponding ``*_by_price_status`` fields distinguish known + prices from unknown-price fallback zeros. """ column = _DIMENSION_TO_COLUMN.get(dimension) if column is None: @@ -1411,17 +1503,30 @@ def rollup( "total_tokens": 0, "cost_amount": Decimal("0"), "currency_code": row.get("currency_code", "USD"), + "cost_amount_by_status": { + status: Decimal("0") for status in MEASUREMENT_STATUSES + }, + "record_count_by_status": {status: 0 for status in MEASUREMENT_STATUSES}, + "cost_amount_by_price_status": { + "known": Decimal("0"), "unknown": Decimal("0") + }, + "record_count_by_price_status": {"known": 0, "unknown": 0}, "_measurement_statuses": set(), }, ) + status = _measurement_status_of(row) + price_status = "known" if row.get("price_known") else "unknown" + row_cost = Decimal(str(row.get("cost_amount", 0))) bucket["record_count"] += 1 bucket["prompt_tokens"] += int(row.get("prompt_tokens", 0)) bucket["completion_tokens"] += int(row.get("completion_tokens", 0)) bucket["total_tokens"] += int(row.get("total_tokens", 0)) - bucket["cost_amount"] += Decimal(str(row.get("cost_amount", 0))) - bucket["_measurement_statuses"].add( - row.get("measurement_status", "unavailable") - ) + bucket["cost_amount"] += row_cost + bucket["cost_amount_by_status"][status] += row_cost + bucket["record_count_by_status"][status] += 1 + bucket["cost_amount_by_price_status"][price_status] += row_cost + bucket["record_count_by_price_status"][price_status] += 1 + bucket["_measurement_statuses"].add(status) for bucket in buckets.values(): statuses = bucket.pop("_measurement_statuses") bucket["measurement_status"] = ( @@ -1438,6 +1543,14 @@ def rollup( ) ) ) + bucket["cost_amount_by_status"] = { + status: float(amount.quantize(Decimal("0.000001"), rounding=ROUND_HALF_UP)) + for status, amount in bucket["cost_amount_by_status"].items() + } + bucket["cost_amount_by_price_status"] = { + status: float(amount.quantize(Decimal("0.000001"), rounding=ROUND_HALF_UP)) + for status, amount in bucket["cost_amount_by_price_status"].items() + } return buckets def report( @@ -1446,7 +1559,13 @@ def report( start: Optional[int] = None, end: Optional[int] = None, ) -> Dict[str, Any]: - """Return a report envelope: per-value rollup plus a grand total.""" + """Return a report envelope: per-value rollup plus a grand total. + + Both ``items`` (from :meth:`rollup`) and ``grand_total`` (from + :meth:`total`) carry the ``cost_amount_by_status``/ + ``record_count_by_status`` measured/estimated/unavailable breakdown + alongside their existing flat ``cost_amount`` totals. + """ buckets = self.rollup(dimension, start, end) items = sorted( buckets.values(), @@ -1465,10 +1584,32 @@ def report( } def total(self, start: Optional[int] = None, end: Optional[int] = None) -> Dict[str, Any]: - """Return grand totals (cost + tokens + record count) over the window.""" + """Return grand totals (cost + tokens + record count) over the window. + + ``cost_amount`` remains the sum over every row regardless of + provenance (unchanged, backward-compatible meaning). It is joined by + ``cost_amount_by_status``/``record_count_by_status``, the same total + broken down by row ``measurement_status`` (``measured``, ``estimated``, + ``unavailable``) so measured, estimated, and unavailable-priced spend + are never opaquely blended into one authoritative-looking number. The + corresponding ``*_by_price_status`` fields expose unknown-price rows. + """ rows = self.store.query(start, end) - cost = sum((Decimal(str(row.get("cost_amount", 0))) for row in rows), Decimal("0")) - statuses = {row.get("measurement_status", "unavailable") for row in rows} + cost_amount_by_status = {status: Decimal("0") for status in MEASUREMENT_STATUSES} + record_count_by_status = {status: 0 for status in MEASUREMENT_STATUSES} + cost_amount_by_price_status = {"known": Decimal("0"), "unknown": Decimal("0")} + record_count_by_price_status = {"known": 0, "unknown": 0} + statuses = set() + for row in rows: + status = _measurement_status_of(row) + price_status = "known" if row.get("price_known") else "unknown" + row_cost = Decimal(str(row.get("cost_amount", 0))) + cost_amount_by_status[status] += row_cost + record_count_by_status[status] += 1 + cost_amount_by_price_status[price_status] += row_cost + record_count_by_price_status[price_status] += 1 + statuses.add(status) + total_cost = sum(cost_amount_by_status.values(), Decimal("0")) measurement_status = ( "unavailable" if "unavailable" in statuses else "estimated" if "estimated" in statuses @@ -1482,9 +1623,19 @@ def total(self, start: Optional[int] = None, end: Optional[int] = None) -> Dict[ "cost_amount": ( None if measurement_status == "unavailable" - else float(cost.quantize(Decimal("0.000001"), rounding=ROUND_HALF_UP)) + else float(total_cost.quantize(Decimal("0.000001"), rounding=ROUND_HALF_UP)) ), "measurement_status": measurement_status, + "cost_amount_by_status": { + status: float(amount.quantize(Decimal("0.000001"), rounding=ROUND_HALF_UP)) + for status, amount in cost_amount_by_status.items() + }, + "record_count_by_status": record_count_by_status, + "cost_amount_by_price_status": { + status: float(amount.quantize(Decimal("0.000001"), rounding=ROUND_HALF_UP)) + for status, amount in cost_amount_by_price_status.items() + }, + "record_count_by_price_status": record_count_by_price_status, } def records(self, start: Optional[int] = None, end: Optional[int] = None) -> List[Dict[str, Any]]: diff --git a/contextual_orchestrator/cost_router.py b/contextual_orchestrator/cost_router.py index c29fe4a84..f4bb29716 100644 --- a/contextual_orchestrator/cost_router.py +++ b/contextual_orchestrator/cost_router.py @@ -115,6 +115,7 @@ def __init__( token_counter=self.token_counter, job_registry=registry ) ) + self._job_registry = registry # job_id -> submitted BatchJob (so poll/retrieve can be driven by id) self._batch_jobs = registry.mapping("batch_jobs", decode=lambda raw: BatchJob(**raw)) # embeddings batch state: job handle + submitted requests + cached doc, @@ -390,6 +391,7 @@ def complete( ) ) currencies = {record.currency_code for record in records} + price_known = all(record.price_known for record in records) provider_response["usage_record_ids"] = [ record.usage_record_id for record in records ] @@ -407,24 +409,23 @@ def complete( provider_response["cost"] = { "cost_amount": ( round(sum(record.cost_amount for record in records), 6) - if len(currencies) == 1 and aggregate_measurement_status != "unavailable" + if price_known and len(currencies) == 1 and aggregate_measurement_status != "unavailable" else None ), "currency_code": next(iter(currencies)) if len(currencies) == 1 else "MIXED", + "price_known": price_known, "measurement_status": aggregate_measurement_status, } - if len(currencies) > 1 and aggregate_measurement_status != "unavailable": + if len(currencies) > 1 and aggregate_measurement_status != "unavailable" and price_known: provider_response["cost"]["currency_components"] = [ { "currency_code": currency, - "cost_amount": round( - sum( - record.cost_amount - for record in records - if record.currency_code == currency - ), - 6, + "cost_amount": ( + round(sum(record.cost_amount for record in records if record.currency_code == currency), 6) + if all(record.price_known for record in records if record.currency_code == currency) + else None ), + "price_known": all(record.price_known for record in records if record.currency_code == currency), } for currency in sorted(currencies) ] @@ -532,36 +533,33 @@ def complete( "total_tokens": sum(item.total_tokens for item in client_usage_records), } currencies = {item.currency_code for item in records} - # Same honesty precedence as the provider_request path above and - # record_stream_usage: a billable race-loser call recorded - # "unavailable" must not be summed into a confident-looking total. statuses = {item.measurement_status for item in records} aggregate_measurement_status = ( "unavailable" if "unavailable" in statuses else "estimated" if "estimated" in statuses else "measured" ) + price_known = all(item.price_known for item in records) result["cost"] = { "cost_amount": ( round(sum(item.cost_amount for item in records), 6) - if len(currencies) == 1 and aggregate_measurement_status != "unavailable" + if price_known and len(currencies) == 1 and aggregate_measurement_status != "unavailable" else None ), "currency_code": next(iter(currencies)) if len(currencies) == 1 else "MIXED", + "price_known": price_known, "measurement_status": aggregate_measurement_status, } - if len(currencies) > 1 and aggregate_measurement_status != "unavailable": + if len(currencies) > 1 and aggregate_measurement_status != "unavailable" and price_known: result["cost"]["currency_components"] = [ { "currency_code": currency, - "cost_amount": round( - sum( - item.cost_amount - for item in records - if item.currency_code == currency - ), - 6, + "cost_amount": ( + round(sum(item.cost_amount for item in records if item.currency_code == currency), 6) + if all(item.price_known for item in records if item.currency_code == currency) + else None ), + "price_known": all(item.price_known for item in records if item.currency_code == currency), } for currency in sorted(currencies) ] @@ -660,6 +658,7 @@ def record_stream_usage( else "measured" ) currencies = {record.currency_code for record in records} + price_known = all(record.price_known for record in records) return { "usage_record_ids": [record.usage_record_id for record in records], "usage": ( @@ -674,11 +673,12 @@ def record_stream_usage( "cost": { "cost_amount": ( round(sum(record.cost_amount for record in records), 6) - if measurement_status == "measured" and len(currencies) == 1 + if measurement_status == "measured" and price_known and len(currencies) == 1 else None ), "currency_code": next(iter(currencies)) if len(currencies) == 1 else "MIXED", "measurement_status": measurement_status, + "price_known": price_known, }, } @@ -700,8 +700,15 @@ def submit_batch( raise BatchModelSelectionError( "no eligible model-group member is available for this batch request" ) from exc + prompt_token_estimates = { + request.custom_id: self.token_counter.count_messages( + request.messages, request.model + ) + for request in prepared_requests + } job = self.batch_backend.submit(prepared_requests, metadata=metadata) job.owner_id = owner_id + job.prompt_token_estimates = prompt_token_estimates self._batch_jobs[job.job_id] = job return job @@ -741,9 +748,25 @@ def retrieve_batch(self, job_id: str, *, owner_id: Optional[str] = None) -> Dict """Retrieve results for a batch owned by ``owner_id`` and record usage.""" job = self._require_job(job_id, owner_id=owner_id) items: List[BatchResultItem] = self.batch_backend.retrieve(job) + prompt_token_estimates = dict(job.prompt_token_estimates) + needs_legacy_lookup = any( + not self._batch_item_usage_valid(item) + and item.custom_id not in prompt_token_estimates + for item in items + ) + request_by_custom_id = ( + self._legacy_batch_requests(job) if needs_legacy_lookup else {} + ) recorded: List[Dict[str, Any]] = [] for item in items: provider_model = self._resolve_batch_provider_model(item) + usage_valid = self._batch_item_usage_valid(item) + if not usage_valid and item.custom_id not in prompt_token_estimates: + original_request = request_by_custom_id.get(item.custom_id) + if original_request is not None: + prompt_token_estimates[item.custom_id] = self.token_counter.count_messages( + original_request.messages, item.model + ) record = self._record_completion( messages=[{"role": "user", "content": ""}], answer=item.answer, @@ -753,21 +776,29 @@ def retrieve_batch(self, job_id: str, *, owner_id: Optional[str] = None) -> Dict model_name=item.model, provider_model=provider_model, workflow_run_id=job.job_id, - prompt_tokens=item.prompt_tokens or None, - completion_tokens=item.completion_tokens or None, + prompt_tokens=( + item.prompt_tokens + if usage_valid + else prompt_token_estimates.get(item.custom_id) + ), + completion_tokens=item.completion_tokens if usage_valid else None, ) recorded.append( { "custom_id": item.custom_id, "answer": item.answer, "usage_record_id": record.usage_record_id, - "cost_amount": record.cost_amount, + "cost_amount": record.cost_amount if record.price_known else None, "currency_code": record.currency_code, + "price_known": record.price_known, "prompt_tokens": record.prompt_tokens, "completion_tokens": record.completion_tokens, "measurement_status": record.measurement_status, } ) + if prompt_token_estimates != job.prompt_token_estimates: + job.prompt_token_estimates = prompt_token_estimates + self._batch_jobs[job.job_id] = job return { "job_id": job_id, "backend": job.backend, @@ -775,6 +806,48 @@ def retrieve_batch(self, job_id: str, *, owner_id: Optional[str] = None) -> Dict "results": recorded, } + @staticmethod + def _batch_item_usage_valid(item: BatchResultItem) -> bool: + """True when a batch result item's provider-reported usage is trustworthy.""" + return ( + item.prompt_tokens >= 0 + and item.completion_tokens >= 0 + and ( + item.usage_valid is True + or ( + item.usage_valid is None + and (item.prompt_tokens > 0 or item.completion_tokens > 0) + ) + ) + ) + + def _legacy_batch_requests(self, job: BatchJob) -> Dict[str, BatchRequest]: + """Read pre-upgrade batch requests for a job with an unestimated custom_id. + + Callers gate this on whether the current retrieval actually needs a + legacy lookup (some item still lacks a stored/computed estimate), not + on whether ``job.prompt_token_estimates`` is merely non-empty -- a + job's estimates can be filled in incrementally across multiple + ``retrieve_batch`` calls, and gating on non-emptiness alone would stop + looking up any custom_id not yet covered by an earlier partial + retrieval. A job never seen by the legacy registry (every job + submitted after this fix) costs one cheap KeyError-guarded miss. + """ + legacy_requests = self._job_registry.mapping( + "batch_requests", decode=lambda raw: BatchRequest(**raw) + ) + try: + requests = legacy_requests[job.job_id] + except KeyError: + return {} + if not isinstance(requests, list): + return {} + return { + request.custom_id: request + for request in requests + if isinstance(request, BatchRequest) + } + def _resolve_batch_provider_model(self, item: BatchResultItem) -> tuple[str, str]: provider = str(item.attribution.get("provider") or item.attribution.get("upstream_api") or "") if not provider: @@ -1072,6 +1145,7 @@ def embeddings_batch_document(self, batch_id: str) -> Dict[str, Any]: embeddings: List[Dict[str, Any]] = [] token_counts: List[int] = [] total_cost_amount = 0.0 + price_known = True currency_code = "USD" for source_index in range(input_count): parts = sorted(parts_by_source.get(source_index, []), key=lambda item: item["part_index"]) @@ -1096,6 +1170,7 @@ def embeddings_batch_document(self, batch_id: str) -> Dict[str, Any]: attribution=attribution, ) total_cost_amount += float(record.cost_amount) + price_known = price_known and record.price_known currency_code = record.currency_code token_counts.append(record.prompt_tokens) embeddings.append( @@ -1121,9 +1196,10 @@ def embeddings_batch_document(self, batch_id: str) -> Dict[str, Any]: "strategy": "token_budgeted_embedding_parts_weighted_average", **part_limits, }, - "cost_amount": round(total_cost_amount, 6), + "cost_amount": round(total_cost_amount, 6) if price_known else None, "currency_code": currency_code, - "cost_micro_usd": int(round(total_cost_amount * 1_000_000)), + "price_known": price_known, + "cost_micro_usd": int(round(total_cost_amount * 1_000_000)) if price_known else None, } self._embedding_documents[batch_id] = document return document @@ -1170,7 +1246,14 @@ def cost_report( start: Optional[int] = None, end: Optional[int] = None, ) -> Dict[str, Any]: - """Return a cost rollup report grouped by ``dimension`` over a window.""" + """Return a cost rollup report grouped by ``dimension`` over a window. + + The envelope's ``items`` and ``grand_total`` each carry a + ``cost_amount_by_status``/``record_count_by_status`` breakdown + (measured/estimated/unavailable) alongside their flat ``cost_amount`` + total, plus known/unknown ``*_by_price_status`` fields — see + :meth:`CostLedger.rollup`. + """ return self.ledger.report(dimension, start, end) diff --git a/contextual_orchestrator/model_discovery.py b/contextual_orchestrator/model_discovery.py index be30fdc7e..6e5656768 100644 --- a/contextual_orchestrator/model_discovery.py +++ b/contextual_orchestrator/model_discovery.py @@ -1893,7 +1893,7 @@ def _discovery_price_key( ): return unknown try: - cost, currency = price_book.compute_cost( + cost, currency, _price_known = price_book.compute_cost( model.provider_name, model.model_id, 1000, diff --git a/tests/test_batch_routing_boundaries.py b/tests/test_batch_routing_boundaries.py index bb04dbbdf..511750457 100644 --- a/tests/test_batch_routing_boundaries.py +++ b/tests/test_batch_routing_boundaries.py @@ -28,9 +28,16 @@ def __init__(self, cost: float) -> None: def compute_cost( self, provider: str, model: str, prompt_tokens: int, completion_tokens: int - ) -> tuple[float, str]: + ) -> tuple[float, str, bool]: self.queries.append((provider, model)) - return self._cost, "USD" + return self._cost, "USD", True + + +class _MixedPriceBook: + def compute_cost( + self, provider: str, model: str, prompt_tokens: int, completion_tokens: int + ) -> tuple[float, str, bool]: + return (2.0, "USD", True) if provider == "known" else (0.0, "USD", False) def test_cheapest_upstream_returns_none_for_no_candidates() -> None: @@ -44,6 +51,13 @@ def test_cheapest_upstream_tie_keeps_input_order() -> None: assert best is first # strict less-than keeps the earlier candidate on ties +def test_cheapest_upstream_excludes_unknown_prices() -> None: + known = {"provider": "known", "model": "priced"} + unknown = {"provider": "unknown", "model": "unpriced"} + assert cheapest_upstream([unknown, known], _MixedPriceBook()) is known + assert cheapest_upstream([unknown], _MixedPriceBook()) is None + + def test_local_backend_rejects_invalid_concurrency() -> None: with pytest.raises(ValueError, match="max_concurrency"): LocalBatchBackend(lambda messages, mode: {}, max_concurrency=0) diff --git a/tests/test_cost_ledger.py b/tests/test_cost_ledger.py index ef454ad53..3702e23fb 100644 --- a/tests/test_cost_ledger.py +++ b/tests/test_cost_ledger.py @@ -139,7 +139,37 @@ def test_unpriced_model_costs_zero_and_still_records() -> None: provider="mystery", model="unpriced", prompt_tokens=100, completion_tokens=50 ) assert record.cost_amount == 0.0 + # The zero must be distinguishable from a real free price: an unpriced + # provider/model is unknown, never fabricated as $0.00. + assert record.price_known is False assert len(ledger.records()) == 1 + assert ledger.total()["record_count_by_price_status"] == {"known": 0, "unknown": 1} + assert ledger.rollup("provider")["mystery"]["record_count_by_price_status"] == { + "known": 0, "unknown": 1, + } + + +def test_priced_model_marks_price_known_true() -> None: + ledger = _priced_ledger() + record = ledger.record_usage( + provider="openai", model="gpt-x", prompt_tokens=100, completion_tokens=50 + ) + assert record.price_known is True + + +def test_compute_cost_returns_price_known_flag() -> None: + """PriceBook.compute_cost's 3-tuple distinguishes known from unknown price.""" + config = InMemoryConfigStore() + price_book = PriceBook(config) + price_book.set_price( + PriceEntry("openai", "gpt-x", prompt_price_per_1k=2.0, completion_price_per_1k=4.0) + ) + + priced = price_book.compute_cost("openai", "gpt-x", 1000, 500) + assert priced == (4.0, "USD", True) + + unpriced = price_book.compute_cost("mystery", "unpriced", 100, 50) + assert unpriced == (0.0, "USD", False) def test_provider_wildcard_price_entry() -> None: @@ -301,6 +331,20 @@ def test_non_blocking_store_records_p2028_like_failure_as_telemetry_only() -> No assert "secret answer" not in repr(sink.events()) +def test_unpriced_usage_telemetry_omits_cost_metric() -> None: + sink = InMemoryUsageTelemetrySink() + ledger = CostLedger(PriceBook(InMemoryConfigStore()), telemetry_sink=sink) + ledger.record_usage( + provider="unknown", + model="unpriced", + prompt_tokens=1, + completion_tokens=1, + ) + event = sink.events()[0] + assert event.attributes["contextual_orchestrator.usage.price_known"] is False + assert "gen_ai.usage.cost" not in event.metrics + + def test_multi_dimensional_rollup_correctness() -> None: ledger = _priced_ledger() ledger.record_usage(provider="openai", model="gpt-x", prompt_tokens=1000, completion_tokens=1000, @@ -392,6 +436,51 @@ def test_report_envelope_sorts_by_cost_desc_and_includes_grand_total() -> None: assert report["grand_total"]["cost_amount"] == 11.0 +def test_rollup_report_total_break_down_cost_by_measurement_status() -> None: + """cost_amount stays a flat sum; the new by-status fields attribute it.""" + ledger = _priced_ledger() + ledger.record_usage( + provider="openai", model="gpt-x", prompt_tokens=1000, completion_tokens=0, + measurement_status="measured", + ) # 2.0, priced + ledger.record_usage( + provider="openai", model="gpt-x", prompt_tokens=500, completion_tokens=0, + measurement_status="estimated", + ) # 1.0, priced + ledger.record_usage( + provider="mystery", model="unpriced", prompt_tokens=1000, completion_tokens=1000, + measurement_status="unavailable", + ) # 0.0, unpriced + + total = ledger.total() + assert total["cost_amount"] is None + assert total["measurement_status"] == "unavailable" + assert total["cost_amount_by_status"] == { + "measured": 2.0, "estimated": 1.0, "unavailable": 0.0, + } + assert total["record_count_by_status"] == { + "measured": 1, "estimated": 1, "unavailable": 1, + } + assert sum(total["cost_amount_by_status"].values()) == 3.0 + + by_provider = ledger.rollup("provider") + assert by_provider["openai"]["cost_amount"] == 3.0 + assert by_provider["openai"]["measurement_status"] == "estimated" + assert by_provider["openai"]["cost_amount_by_status"] == { + "measured": 2.0, "estimated": 1.0, "unavailable": 0.0, + } + assert by_provider["openai"]["record_count_by_status"] == { + "measured": 1, "estimated": 1, "unavailable": 0, + } + assert by_provider["mystery"]["cost_amount"] is None + assert by_provider["mystery"]["measurement_status"] == "unavailable" + assert by_provider["mystery"]["cost_amount_by_status"]["unavailable"] == 0.0 + assert by_provider["mystery"]["record_count_by_status"]["unavailable"] == 1 + + report = ledger.report("provider") + assert report["grand_total"] == total + + def test_sql_ledger_store_on_sqlite_creates_objects_and_rolls_up() -> None: conn = sqlite3.connect(":memory:") store = SqlLedgerStore(conn, paramstyle="qmark") @@ -414,6 +503,23 @@ def test_sql_ledger_store_on_sqlite_creates_objects_and_rolls_up() -> None: assert by_company["acme"]["cost_amount"] == 15.0 # 6 + 9 +def test_sql_ledger_persists_price_known_flag() -> None: + """price_known round-trips through the SQL store's satellite table.""" + conn = sqlite3.connect(":memory:") + store = SqlLedgerStore(conn, paramstyle="qmark") + ledger = _priced_ledger(store=store) + ledger.record_usage(provider="openai", model="gpt-x", prompt_tokens=1000, completion_tokens=0) + ledger.record_usage(provider="mystery", model="unpriced", prompt_tokens=1000, completion_tokens=0) + + assert conn.execute( + "SELECT name FROM sqlite_master WHERE type='table' AND name='usage_price_knowledge'" + ).fetchone() == ("usage_price_knowledge",) + + rows = {row["provider_name"]: row for row in store.query()} + assert bool(rows["openai"]["price_known"]) is True + assert bool(rows["mystery"]["price_known"]) is False + + def test_table_columns_honors_its_argument_on_qmark() -> None: """The qmark branch must inspect the requested table, not a hardcoded one.""" conn = sqlite3.connect(":memory:") @@ -581,6 +687,10 @@ def test_sql_ledger_migrates_flattened_usage_rows() -> None: assert row["service_name"] == "unattributed" assert row["team_name"] == "unattributed" assert row["company_name"] == "acme" + # A flattened legacy row predates the price_known signal: whether its + # stored cost reflects a real price is unknown, so migration marks it + # unknown rather than assuming it was known. + assert bool(row["price_known"]) is False def test_sql_ledger_maps_null_legacy_attribution_to_unattributed() -> None: @@ -861,6 +971,8 @@ def test_orphaned_legacy_generation_is_adopted_and_dropped() -> None: rows = store.query(None, None) assert len(rows) == 1 assert rows[0]["usage_record_id"] == "usage_guard_t1" + # Adopted from a pre-price_known legacy table: unknown, not assumed known. + assert bool(rows[0]["price_known"]) is False legacy_tables = connection.execute( "SELECT COUNT(*) FROM sqlite_master WHERE name='llm_usage_records_legacy'" ).fetchone()[0] @@ -904,6 +1016,35 @@ def test_append_preserves_caller_owned_sqlite_transaction() -> None: assert store.query(None, None) == [] +def test_duplicate_append_does_not_rewrite_price_provenance() -> None: + """Devin review (#956): a rejected duplicate parent insert must stay a + true no-op. Previously the satellite usage_measurements/ + usage_price_knowledge/attribution writes ran unconditionally even when + the parent row already existed, so a retry with today's price state + could backfill a historical price-unknown row's provenance from an + unrelated later call instead of what actually priced that original + spend (the exact gap an upgrade-migrated pre-price-knowledge row is + supposed to preserve). + """ + connection = sqlite3.connect(":memory:", isolation_level=None) + store = SqlLedgerStore(connection, paramstyle="qmark") + + accepted_first = store.append( + _guard_usage_record(measurement_status="unavailable", price_known=False) + ) + assert accepted_first is True + + accepted_retry = store.append( + _guard_usage_record(measurement_status="measured", price_known=True) + ) + assert accepted_retry is False + + rows = store.query(None, None) + assert len(rows) == 1 + assert rows[0]["measurement_status"] == "unavailable" + assert not rows[0]["price_known"] + + def test_sql_ledger_rejects_unknown_parameter_style() -> None: try: SqlLedgerStore(sqlite3.connect(":memory:"), paramstyle="named") @@ -947,6 +1088,8 @@ def test_ledger_table_names_follow_two_word_snake_case() -> None: "llm_price_entries", "cost_attribution_values", "usage_record_attributions", + "usage_measurements", + "usage_price_knowledge", ): assert is_two_word_snake_case(name) diff --git a/tests/test_cost_ledger_boundaries.py b/tests/test_cost_ledger_boundaries.py index 77fcba4d4..a4f0efbad 100644 --- a/tests/test_cost_ledger_boundaries.py +++ b/tests/test_cost_ledger_boundaries.py @@ -268,9 +268,10 @@ def test_corrupt_price_component_text_is_rejected_not_parsed() -> None: ) book = PriceBook(config) assert book.get_price("openai", "corrupt-model") is None - cost, currency = book.compute_cost("openai", "corrupt-model", 1000, 1000) + cost, currency, price_known = book.compute_cost("openai", "corrupt-model", 1000, 1000) assert cost == 0.0 assert currency == "USD" + assert price_known is False def test_working_background_store_marks_records_stored() -> None: diff --git a/tests/test_cost_router.py b/tests/test_cost_router.py index 63eafa062..6c3d162ef 100644 --- a/tests/test_cost_router.py +++ b/tests/test_cost_router.py @@ -572,6 +572,7 @@ def run(*_args, **_kwargs): assert result["cost"] == { "cost_amount": None, "currency_code": "MIXED", + "price_known": True, "measurement_status": "unavailable", } @@ -696,8 +697,8 @@ def test_structured_mixed_currency_costs_are_never_implicitly_converted() -> Non assert result["cost"]["cost_amount"] is None assert result["cost"]["currency_code"] == "MIXED" assert result["cost"]["currency_components"] == [ - {"currency_code": "EUR", "cost_amount": 2.0}, - {"currency_code": "USD", "cost_amount": 1.0}, + {"currency_code": "EUR", "cost_amount": 2.0, "price_known": True}, + {"currency_code": "USD", "cost_amount": 1.0, "price_known": True}, ] assert "approved exchange-rate source" in result["cost"]["customer_action"] diff --git a/tests/test_cost_router_boundaries.py b/tests/test_cost_router_boundaries.py index d09d75ab5..aa778f865 100644 --- a/tests/test_cost_router_boundaries.py +++ b/tests/test_cost_router_boundaries.py @@ -15,9 +15,11 @@ ) from contextual_orchestrator.batch_routing import ( BatchJob, + BatchRequest, BatchResultItem, EmbeddingBatchResultItem, ) +from contextual_orchestrator.batch_job_registry import JobRegistryFactory from contextual_orchestrator.cost_router import ( CostRoutingCoordinator as Coordinator, ) @@ -96,6 +98,253 @@ def test_stream_usage_aggregates_trace_steps_without_text_estimates() -> None: assert len(coordinator.ledger.records()) == 2 +def test_unpriced_stream_usage_omits_cost_total() -> None: + result = _coordinator().record_stream_usage( + result={ + "workflow_run_id": "unpriced-stream", + "mode": "route", + "trace": [{"agent_id": "mock_worker", "usage": {"prompt_tokens": 1, "completion_tokens": 1}}], + }, + attribution=None, + model_name="mock-a", + ) + assert result["cost"]["price_known"] is False + assert result["cost"]["cost_amount"] is None + + +def test_unpriced_sync_cost_omits_total() -> None: + result = _coordinator().complete([{"role": "user", "content": "hello"}]) + assert result["cost"]["price_known"] is False + assert result["cost"]["cost_amount"] is None + + +def test_unpriced_provider_request_cost_omits_total() -> None: + result = _coordinator().complete( + [{"role": "user", "content": "return json"}], + provider_request={ + "model": "mock-a", + "messages": [{"role": "user", "content": "return json"}], + "response_format": {"type": "json_object"}, + }, + ) + assert result["cost"]["price_known"] is False + assert result["cost"]["cost_amount"] is None + + +def test_unpriced_batch_item_omits_cost() -> None: + coordinator = _coordinator() + job = coordinator.submit_batch( + [BatchRequest(messages=[{"role": "user", "content": "batch"}], model="mock-a")] + ) + item = coordinator.retrieve_batch(job.job_id)["results"][0] + assert item["price_known"] is False + assert item["cost_amount"] is None + + +def test_provider_confirmed_zero_usage_stays_measured_and_price_known() -> None: + class ZeroUsageBackend: + name = "zero-usage" + + def submit(self, requests, metadata=None): # type: ignore[no-untyped-def] + del requests, metadata + return BatchJob("zero-usage-job", self.name, request_count=1) + + def retrieve(self, job): # type: ignore[no-untyped-def] + del job + return [BatchResultItem( + "zero-result", "", prompt_tokens=0, completion_tokens=0, + model="unpriced-model", usage_valid=True, + )] + + coordinator = _coordinator(batch_backend=ZeroUsageBackend()) + job = coordinator.submit_batch([BatchRequest( + messages=[{"role": "user", "content": "ignored"}], model="unpriced-model" + )]) + + item = coordinator.retrieve_batch(job.job_id)["results"][0] + + assert item["prompt_tokens"] == item["completion_tokens"] == 0 + assert item["measurement_status"] == "measured" + assert item["price_known"] is True + assert item["cost_amount"] == 0.0 + + +def test_invalid_batch_usage_estimates_prompt_tokens_from_original_request() -> None: + """usage_valid=False must estimate from the real submitted prompt. + + The fallback count is computed before submission and stored as safe + metadata on the durable job record. Raw prompts are never copied into the + shared registry, and the accepted job has only one publication write. + """ + class InvalidUsageBackend: + name = "invalid-usage" + + def submit(self, requests, metadata=None): # type: ignore[no-untyped-def] + del metadata + self._custom_id = requests[0].custom_id + return BatchJob("invalid-usage-job", self.name, request_count=1) + + def retrieve(self, job): # type: ignore[no-untyped-def] + del job + return [BatchResultItem( + self._custom_id, "short answer", prompt_tokens=-1, completion_tokens=2, + model="mock-a", usage_valid=False, + )] + + coordinator = _coordinator(batch_backend=InvalidUsageBackend()) + large_prompt = "word " * 500 + job = coordinator.submit_batch([ + BatchRequest(messages=[{"role": "user", "content": large_prompt}], model="mock-a") + ]) + + stored_job = coordinator._batch_jobs[job.job_id] # noqa: SLF001 + assert stored_job.prompt_token_estimates == job.prompt_token_estimates + assert large_prompt not in repr(stored_job) + + item = coordinator.retrieve_batch(job.job_id)["results"][0] + + assert item["measurement_status"] == "estimated" + # A hardcoded empty placeholder would estimate near-zero prompt tokens + # regardless of the real prompt's size; the actual submitted prompt must + # drive the estimate instead. + assert item["prompt_tokens"] > 100 + + +def test_invalid_batch_usage_uses_legacy_request_registry_when_job_predates_prompt_estimates() -> None: + """Pre-upgrade jobs keep old prompt fallback compatibility without new writes.""" + class LegacyRegistry(JobRegistryFactory): + def __init__(self) -> None: + super().__init__() + self._mappings: dict[str, Any] = {} + + def mapping(self, name, *, decode=None): # type: ignore[no-untyped-def] + if name not in self._mappings: + self._mappings[name] = {} + return self._mappings[name] + + class InvalidUsageBackend: + name = "invalid-usage" + + def retrieve(self, job): # type: ignore[no-untyped-def] + del job + return [BatchResultItem( + "legacy-request", "short answer", prompt_tokens=-1, completion_tokens=2, + model="mock-a", usage_valid=False, + )] + + registry = LegacyRegistry() + coordinator = _coordinator(batch_backend=InvalidUsageBackend(), job_registry=registry) + coordinator._batch_jobs["legacy-job"] = BatchJob( # noqa: SLF001 + "legacy-job", "invalid-usage", request_count=1 + ) + registry.mapping("batch_requests", decode=lambda raw: BatchRequest(**raw))["legacy-job"] = [ + BatchRequest( + messages=[{"role": "user", "content": "word " * 500}], + model="mock-a", + custom_id="legacy-request", + ) + ] + + item = coordinator.retrieve_batch("legacy-job")["results"][0] + + assert item["measurement_status"] == "estimated" + assert item["prompt_tokens"] > 100 + assert coordinator._batch_jobs["legacy-job"].prompt_token_estimates == { # noqa: SLF001 + "legacy-request": item["prompt_tokens"] + } + + +def test_legacy_batch_requests_stay_available_across_partial_retrievals() -> None: + """A partial retrieval must not stop a later one from reaching legacy fallback. + + (Devin review on #956) Once one legacy custom_id picks up a stored + estimate, ``job.prompt_token_estimates`` becomes non-empty -- gating the + legacy registry lookup on mere non-emptiness would silently stop looking + up every other still-unestimated custom_id from that point on, even + though the legacy request registry still holds their original prompts. + """ + class LegacyRegistry(JobRegistryFactory): + def __init__(self) -> None: + super().__init__() + self._mappings: dict[str, Any] = {} + + def mapping(self, name, *, decode=None): # type: ignore[no-untyped-def] + if name not in self._mappings: + self._mappings[name] = {} + return self._mappings[name] + + class TwoPassInvalidUsageBackend: + name = "invalid-usage" + + def __init__(self) -> None: + self.calls = 0 + + def retrieve(self, job): # type: ignore[no-untyped-def] + del job + self.calls += 1 + custom_id = "legacy-request-one" if self.calls == 1 else "legacy-request-two" + return [BatchResultItem( + custom_id, "short answer", prompt_tokens=-1, completion_tokens=2, + model="mock-a", usage_valid=False, + )] + + registry = LegacyRegistry() + backend = TwoPassInvalidUsageBackend() + coordinator = _coordinator(batch_backend=backend, job_registry=registry) + coordinator._batch_jobs["legacy-job"] = BatchJob( # noqa: SLF001 + "legacy-job", "invalid-usage", request_count=2 + ) + registry.mapping("batch_requests", decode=lambda raw: BatchRequest(**raw))["legacy-job"] = [ + BatchRequest( + messages=[{"role": "user", "content": "word " * 500}], + model="mock-a", + custom_id="legacy-request-one", + ), + BatchRequest( + messages=[{"role": "user", "content": "word " * 700}], + model="mock-a", + custom_id="legacy-request-two", + ), + ] + + first = coordinator.retrieve_batch("legacy-job")["results"][0] + assert first["prompt_tokens"] > 100 # first partial retrieval: legacy lookup works + + second = coordinator.retrieve_batch("legacy-job")["results"][0] + # Pre-fix: the first retrieval already left prompt_token_estimates + # non-empty, so this second retrieval's legacy lookup was skipped + # entirely and this custom_id's real prompt was never found. + assert second["custom_id"] == "legacy-request-two" + assert second["prompt_tokens"] > 100 + + assert coordinator._batch_jobs["legacy-job"].prompt_token_estimates == { # noqa: SLF001 + "legacy-request-one": first["prompt_tokens"], + "legacy-request-two": second["prompt_tokens"], + } + + +def test_batch_prompt_fallback_has_no_separate_registry_publication() -> None: + """Accepted jobs publish their safe fallback metadata in one job record.""" + class RecordingRegistry(JobRegistryFactory): + def __init__(self) -> None: + super().__init__() + self.names: list[str] = [] + + def mapping(self, name, *, decode=None): # type: ignore[no-untyped-def] + self.names.append(name) + return super().mapping(name, decode=decode) + + registry = RecordingRegistry() + coordinator = _coordinator(job_registry=registry) + + coordinator.submit_batch([ + BatchRequest(messages=[{"role": "user", "content": "private prompt"}]) + ]) + + assert "batch_jobs" in registry.names + assert "batch_requests" not in registry.names + + def test_complete_rejects_non_boolean_cache_bypass() -> None: coordinator = _coordinator() with pytest.raises(TypeError, match="cache_bypass"): @@ -327,6 +576,13 @@ def test_embeddings_document_is_idempotent_after_completion() -> None: assert backend.polled.count(job.job_id) == 1 +def test_unpriced_embeddings_document_omits_cost() -> None: + document = _coordinator().complete_embeddings_batch(["unpriced embedding"]) + assert document["price_known"] is False + assert document["cost_amount"] is None + assert document["cost_micro_usd"] is None + + def test_embeddings_document_requires_known_batch() -> None: coordinator = _coordinator() with pytest.raises(KeyError, match="embeddings batch job"): @@ -442,7 +698,8 @@ def test_complete_embeddings_batch_round_trips_locally() -> None: ) assert document["status"] == "completed" assert document["embeddings"][0]["index"] == 0 - assert document["cost_micro_usd"] >= 0 + assert document["cost_micro_usd"] is None + assert document["price_known"] is False diff --git a/tests/test_distributed_cache_truth_and_isolation.py b/tests/test_distributed_cache_truth_and_isolation.py index 13dbcf1ff..1d45ebe8b 100644 --- a/tests/test_distributed_cache_truth_and_isolation.py +++ b/tests/test_distributed_cache_truth_and_isolation.py @@ -111,6 +111,7 @@ def test_cache_hit_records_zero_provider_usage_instead_of_rebilling_inference() assert second["cost"] == { "cost_amount": 0.0, "currency_code": "USD", + "price_known": True, "measurement_status": "measured", }