diff --git a/ee/src/shim_enterprise/application.py b/ee/src/shim_enterprise/application.py index dbae432..a6812fc 100644 --- a/ee/src/shim_enterprise/application.py +++ b/ee/src/shim_enterprise/application.py @@ -135,6 +135,7 @@ async def _lifespan(application: FastAPI) -> AsyncIterator[None]: try: yield finally: + await application.state.gateway_service.kernel.postprocessor.drain() await http_client.aclose() await cache.close() await engine.dispose() diff --git a/ee/src/shim_enterprise/gateway/pipeline/quota_reservation.py b/ee/src/shim_enterprise/gateway/pipeline/quota_reservation.py index 9c11089..f62544c 100644 --- a/ee/src/shim_enterprise/gateway/pipeline/quota_reservation.py +++ b/ee/src/shim_enterprise/gateway/pipeline/quota_reservation.py @@ -97,7 +97,7 @@ async def quota( select(TierDefinition) .where(TierDefinition.slug == api_key.tier) .execution_options(populate_existing=True) - .with_for_update() + .with_for_update(read=True) ) tier = (await session.execute(tier_statement)).scalar_one_or_none() if tier is None: diff --git a/ee/tests/architecture/test_manual_test_dashboard.py b/ee/tests/architecture/test_manual_test_dashboard.py index e7db707..3f2bdde 100644 --- a/ee/tests/architecture/test_manual_test_dashboard.py +++ b/ee/tests/architecture/test_manual_test_dashboard.py @@ -86,7 +86,7 @@ async def test_api_lifespan_does_not_start_continuous_reconciliation( monkeypatch, ) -> None: cache = SimpleNamespace(close=AsyncMock()) - kernel = SimpleNamespace() + kernel = SimpleNamespace(postprocessor=SimpleNamespace(drain=AsyncMock())) create_gateway_kernel = Mock(return_value=kernel) engine = SimpleNamespace(dispose=AsyncMock()) connect_cache = AsyncMock() @@ -110,6 +110,7 @@ async def test_api_lifespan_does_not_start_continuous_reconciliation( assert create_gateway_kernel.call_args.args[1].is_closed connect_cache.assert_awaited_once_with(cache) create_task.assert_not_called() + kernel.postprocessor.drain.assert_awaited_once() cache.close.assert_awaited_once() engine.dispose.assert_awaited_once() shutdown_tracing.assert_called_once() diff --git a/ee/tests/gateway/kernel/test_accounting_coordinator.py b/ee/tests/gateway/kernel/test_accounting_coordinator.py index 0e87868..615818c 100644 --- a/ee/tests/gateway/kernel/test_accounting_coordinator.py +++ b/ee/tests/gateway/kernel/test_accounting_coordinator.py @@ -1506,3 +1506,148 @@ async def test_audit_intent_outbox_reference_is_tenant_scoped(db) -> None: }, ) assert "fk_audit_intent_org_outbox_event" in str(reference_error.value.orig) + + +@pytest.mark.asyncio +async def test_quota_policy_uses_shared_tier_lock_and_exclusive_key_lock(): + from sqlalchemy.dialects import postgresql + from shim_enterprise.gateway.pipeline.quota_reservation import ( + AccountingPolicyLoader, + ) + + session = SimpleNamespace( + execute=AsyncMock( + side_effect=[ + SimpleNamespace( + scalar_one_or_none=lambda: SimpleNamespace(tier="free") + ), + SimpleNamespace( + scalar_one_or_none=lambda: SimpleNamespace( + slug="free", + daily_request_limit=10, + monthly_request_limit=100, + monthly_token_limit=1000, + ) + ), + ] + ) + ) + policy = await AccountingPolicyLoader().quota(session, _prepared()) + statements = [ + str(call.args[0].compile(dialect=postgresql.dialect())) + for call in session.execute.await_args_list + ] + assert statements[0].endswith("FOR UPDATE") + assert statements[1].endswith("FOR SHARE") + assert policy.daily_request_limit == 10 + + +@pytest.mark.asyncio +async def test_shared_tier_allows_independent_reservations_but_fences_edits( + async_engine, +): + from sqlalchemy import update + from shim_enterprise.billing.ledger import QuotaLimitExceeded + from shim_enterprise.gateway.pipeline.quota_reservation import ( + AccountingPolicyLoader, + ) + from shim_enterprise.tenants.models import TierDefinition + + factory = async_sessionmaker(async_engine, expire_on_commit=False) + slug = f"lock-test-{uuid4().hex}" + async with factory.begin() as setup: + setup.add( + TierDefinition( + slug=slug, + name="Lock test", + rate_limit_rpm=60, + rate_limit_tpm=1000, + daily_request_limit=1, + monthly_request_limit=1, + monthly_token_limit=1000, + ) + ) + await setup.flush() + tenants = [await _create_tenant(setup, label) for label in ("lock-a", "lock-b")] + await setup.execute( + update(ApiKey) + .where(ApiKey.id.in_([key for _, _, key in tenants])) + .values(tier=slug) + ) + + async def reserve(session, tenant): + prepared = _prepared() + prepared.tenant_id, _, prepared.api_key_id = tenant + policy = await AccountingPolicyLoader().quota(session, prepared) + now = datetime.now(timezone.utc) + await DurableAccountingRepository().reserve_quota( + session, + QuotaReservationCommand( + tenant_id=prepared.tenant_id, + api_key_id=prepared.api_key_id, + request_id=prepared.request_id, + requested_model=prepared.model, + source_endpoint="chat.completions", + started_at=now, + reconciliation_due_at=now + timedelta(minutes=2), + estimated_input_tokens=1, + maximum_output_tokens=1, + policy=policy, + ), + ) + return policy + + try: + async with factory() as first, factory() as second, factory() as contender: + await reserve(first, tenants[0]) + await second.execute(text("SET LOCAL lock_timeout = '500ms'")) + # This must succeed while the first tenant still holds its tier lock. + snapshot = await reserve(second, tenants[1]) + assert snapshot.daily_request_limit == 1 + + await contender.execute(text("SET LOCAL lock_timeout = '100ms'")) + with pytest.raises(DBAPIError) as same_key: + await reserve(contender, tenants[0]) + assert same_key.value.orig.sqlstate == "55P03" + await contender.rollback() + await first.commit() + with pytest.raises(QuotaLimitExceeded): + await reserve(contender, tenants[0]) + await contender.rollback() + + tier_update = ( + update(TierDefinition) + .where(TierDefinition.slug == slug) + .values(daily_request_limit=2, monthly_request_limit=2) + ) + await contender.execute(text("SET LOCAL lock_timeout = '100ms'")) + with pytest.raises(DBAPIError) as policy_edit: + await contender.execute(tier_update) + assert policy_edit.value.orig.sqlstate == "55P03" + await contender.rollback() + await second.rollback() + await contender.execute(tier_update) + await contender.commit() + revised = await reserve(contender, tenants[1]) + assert (revised.daily_request_limit, revised.monthly_request_limit) == ( + 2, + 2, + ) + assert revised.version != snapshot.version + finally: + async with factory.begin() as cleanup: + ids = [organization for organization, _, _ in tenants] + for model in ( + UsageLedger, + RequestLifecycle, + QuotaPeriodUsage, + ApiKey, + User, + ): + await cleanup.execute( + delete(model).where(model.organization_id.in_(ids)) + ) + await cleanup.execute(delete(Organization).where(Organization.id.in_(ids))) + await cleanup.execute( + delete(TierDefinition).where(TierDefinition.slug == slug) + ) diff --git a/src/shim/application.py b/src/shim/application.py index 9745bc3..decb59f 100644 --- a/src/shim/application.py +++ b/src/shim/application.py @@ -121,6 +121,8 @@ async def lifespan(application: FastAPI) -> AsyncIterator[None]: try: yield finally: + await application.state.gateway_service.kernel.postprocessor.drain() + await usage.aclose() if owns_http_client: await client.aclose() diff --git a/src/shim/gateway/admission.py b/src/shim/gateway/admission.py index 2c70411..d3cc4de 100644 --- a/src/shim/gateway/admission.py +++ b/src/shim/gateway/admission.py @@ -6,6 +6,8 @@ from collections.abc import Callable, Hashable from dataclasses import dataclass import hashlib +import heapq +from itertools import count import time from typing import Literal, Protocol @@ -39,6 +41,8 @@ def __init__( self.max_entries = max_entries self.clock = clock self.windows: OrderedDict[Hashable, tuple[int, float]] = OrderedDict() + self._expirations: list[tuple[float, int, Hashable]] = [] + self._sequence = count() def increment(self, key: Hashable, *, amount: int, window_seconds: int) -> int: now = self.clock() @@ -52,15 +56,25 @@ def increment(self, key: Hashable, *, amount: int, window_seconds: int) -> int: while len(self.windows) >= self.max_entries: self.windows.popitem(last=False) count = amount - self.windows[key] = (count, now + window_seconds) + expires_at = now + window_seconds + self.windows[key] = (count, expires_at) + heapq.heappush(self._expirations, (expires_at, next(self._sequence), key)) + if len(self._expirations) > 2 * self.max_entries: + self._expirations = [ + (expiry, next(self._sequence), item) + for item, (_, expiry) in self.windows.items() + ] + heapq.heapify(self._expirations) return count count = current[0] + amount self.windows[key] = (count, current[1]) return count def _discard_expired(self, now: float) -> None: - for key, (_, expires_at) in tuple(self.windows.items()): - if expires_at <= now: + while self._expirations and self._expirations[0][0] <= now: + expires_at, _, key = heapq.heappop(self._expirations) + current = self.windows.get(key) + if current is not None and current[1] == expires_at: del self.windows[key] diff --git a/src/shim/gateway/pipeline/admission.py b/src/shim/gateway/pipeline/admission.py index 9f31ac4..17b9d7b 100644 --- a/src/shim/gateway/pipeline/admission.py +++ b/src/shim/gateway/pipeline/admission.py @@ -131,8 +131,7 @@ async def run(self, value: PreparedInference) -> PreparedInference: }, ) input_tokens = _estimate_input_tokens(payload) - candidate_count = _candidate_count(payload) - output_tokens = per_candidate_output_tokens * candidate_count + output_tokens = per_candidate_output_tokens * candidate_count(value) tier = value.context.tier_policy key_hash = value.policy.rate_limit_key_hash if tier.rate_limit_rpm is not None and not await self.rate_limiter.allow( @@ -206,13 +205,16 @@ def _estimate_input_tokens(payload: Mapping[str, object]) -> int: return max(1, len(serialized.encode("utf-8", errors="backslashreplace"))) -def _candidate_count(payload: Mapping[str, object]) -> int: - generation_config = payload.get("generationConfig") - candidate = ( - generation_config.get("candidateCount") - if isinstance(generation_config, Mapping) - else payload.get("n", 1) - ) +def candidate_count(prepared: PreparedInference) -> int: + if prepared.provider == "google": + config = prepared.payload.get("generationConfig") + candidate = ( + config.get("candidateCount", 1) if isinstance(config, Mapping) else 1 + ) + elif prepared.provider == "openai" and prepared.protocol == "chat": + candidate = prepared.payload.get("n", 1) + else: + candidate = 1 count = ( candidate if isinstance(candidate, int) and not isinstance(candidate, bool) diff --git a/src/shim/gateway/pipeline/anthropic_execution.py b/src/shim/gateway/pipeline/anthropic_execution.py index 937ed5b..d3f47c5 100644 --- a/src/shim/gateway/pipeline/anthropic_execution.py +++ b/src/shim/gateway/pipeline/anthropic_execution.py @@ -147,7 +147,8 @@ async def close_stream() -> None: return state["closed"] = True try: - await result.close() + async with asyncio.timeout(5): + await result.close() except Exception: pass finally: @@ -334,6 +335,7 @@ def _error_event(error: ProviderCallError) -> bytes: "type": "error", "error": { "type": "api_error", + "code": error.error_code, "message": message, }, } diff --git a/src/shim/gateway/pipeline/google_execution.py b/src/shim/gateway/pipeline/google_execution.py index 1969eef..b9edfc1 100644 --- a/src/shim/gateway/pipeline/google_execution.py +++ b/src/shim/gateway/pipeline/google_execution.py @@ -19,7 +19,7 @@ ProviderStream, ) from shim.gateway.streaming.sse import encode_data -from shim.privacy.deanonymizer import _split_placeholder_prefix +from shim.privacy.deanonymizer import restore_fragment from shim.privacy.pii_scrubber import PIIScrubberService from shim.secrets.credentials import ProviderCredentialResolver @@ -206,7 +206,7 @@ async def _stream( except (asyncio.CancelledError, GeneratorExit): raise except Exception as exc: - if state["recorded"]: + if state["recorded"] and not isinstance(exc, ValueError): return await self._record_error(exc) state["recorded"] = True @@ -289,11 +289,9 @@ def _restore_value( return value def _restore_fragment(self, key: tuple[object, ...], fragment: str) -> str: - text = self._buffers.pop(key, "") + fragment - ready, carry = _split_placeholder_prefix(text, self._verification_map) - if carry: - self._buffers[key] = carry - return self._scrubber.deanonymize(ready, self._verification_map) + return restore_fragment( + self._buffers, key, fragment, self._verification_map, self._scrubber + ) def _flush_candidate( self, @@ -481,7 +479,8 @@ def _stream_error(error: ProviderCallError) -> bytes: async def _close_client(client: genai.Client) -> None: try: - await client.aio.aclose() + async with asyncio.timeout(5): + await client.aio.aclose() except Exception: pass try: @@ -492,6 +491,7 @@ async def _close_client(client: genai.Client) -> None: async def _close_stream(stream) -> None: try: - await stream.aclose() + async with asyncio.timeout(5): + await stream.aclose() except Exception: pass diff --git a/src/shim/gateway/pipeline/openai_execution.py b/src/shim/gateway/pipeline/openai_execution.py index c8df21e..07b9f99 100644 --- a/src/shim/gateway/pipeline/openai_execution.py +++ b/src/shim/gateway/pipeline/openai_execution.py @@ -131,7 +131,8 @@ async def close_stream() -> None: return state["closed"] = True try: - await result.close() + async with asyncio.timeout(5): + await result.close() except Exception: pass finally: @@ -326,7 +327,7 @@ async def _chat_stream( except (asyncio.CancelledError, GeneratorExit): raise except Exception as exc: - if state["recorded"]: + if state["recorded"] and not isinstance(exc, ValueError): yield b"data: [DONE]\n\n" return await self._record_error(exc) diff --git a/src/shim/gateway/pipeline/postprocess.py b/src/shim/gateway/pipeline/postprocess.py index 0eee3ee..fe0ab2a 100644 --- a/src/shim/gateway/pipeline/postprocess.py +++ b/src/shim/gateway/pipeline/postprocess.py @@ -2,7 +2,7 @@ from __future__ import annotations -import json +import asyncio from collections.abc import Mapping from datetime import datetime, timezone from time import perf_counter @@ -12,6 +12,7 @@ from shim.billing.pricing import DEFAULT_PRICE_BOOK, compute_cost_usd from shim.gateway.kernel.result import PreparedInference, UNSPECIFIED_PROVIDER_MODEL +from shim.gateway.pipeline.admission import candidate_count from shim.gateway.pipeline.provider_execution import ProviderNonStream, ProviderStream from shim.gateway.streaming import ( StreamFinalization, @@ -52,10 +53,19 @@ def __init__( heartbeat_interval_seconds: float, output_hash_salt: str | None, ) -> None: + self._finalization_tasks: set[asyncio.Task[Any]] = set() self.usage = usage self.heartbeat_interval_seconds = heartbeat_interval_seconds self.output_hash_salt = output_hash_salt + async def drain(self, timeout_seconds: float = 5.0) -> None: + if self._finalization_tasks: + _, pending = await asyncio.wait( + tuple(self._finalization_tasks), timeout=timeout_seconds + ) + for task in pending: + task.cancel() + async def finalize( self, prepared: PreparedInference, @@ -96,7 +106,7 @@ async def finalize( lifecycle_status = _lifecycle_status( response.payload, provider=provider, - expected_candidates=_expected_candidates(prepared), + expected_candidates=candidate_count(prepared), ) response_model = response.payload.get("model") settlement_model = ( @@ -125,11 +135,6 @@ async def finalize( ), ).inc() PROVIDER_LATENCY_MS.labels(**labels).observe(response.latency_ms) - body = json.dumps( - response.payload, - ensure_ascii=False, - separators=(",", ":"), - ) gateway_response = JSONResponse( content=response.payload, headers=_gateway_headers(prepared, response.request_id), @@ -152,7 +157,9 @@ async def finalize( ), estimated=not fully_actual, output_hash=( - content_ref(self.output_hash_salt, body) + content_ref( + self.output_hash_salt, bytes(gateway_response.body).decode() + ) if self.output_hash_salt is not None else None ), @@ -211,7 +218,7 @@ def observe_terminal(terminal_status: str) -> None: provider=str(prepared.provider), requested_model=prepared.model, prompt_tokens_estimated=prepared.admission.estimated_input_tokens, - expected_candidates=_expected_candidates(prepared), + expected_candidates=candidate_count(prepared), output_hash_salt=self.output_hash_salt, ), finalizer=finalize_stream, @@ -219,6 +226,7 @@ def observe_terminal(terminal_status: str) -> None: stream_heartbeat_recorder=record_stream_heartbeat, heartbeat_interval_seconds=self.heartbeat_interval_seconds, terminal_observer=observe_terminal, + finalization_tasks=self._finalization_tasks, ) @@ -341,21 +349,6 @@ def _sum_counts(*values: int | None) -> int | None: return sum(present) if present else None -def _expected_candidates(prepared: PreparedInference) -> int: - if prepared.provider == "openai" and prepared.protocol == "chat": - count = prepared.payload.get("n", 1) - elif prepared.provider == "google": - config = prepared.payload.get("generationConfig") - count = config.get("candidateCount", 1) if isinstance(config, Mapping) else 1 - else: - count = 1 - return ( - count - if isinstance(count, int) and not isinstance(count, bool) and count > 0 - else 1 - ) - - def _gateway_headers( prepared: PreparedInference, upstream_request_id: str | None, diff --git a/src/shim/gateway/streaming/meter.py b/src/shim/gateway/streaming/meter.py index 1e78085..95842e4 100644 --- a/src/shim/gateway/streaming/meter.py +++ b/src/shim/gateway/streaming/meter.py @@ -218,7 +218,19 @@ def _capture_terminal_hint(self, event_type: str, payload: dict[str, Any]) -> No if payload.get("code") == "PRIVACY_STATE_UNAVAILABLE": self._set_terminal_hint("internal_error") return - if "timeout" in lowered: + error = payload.get("error") + error_fields = error if isinstance(error, dict) else {} + codes = { + str(payload.get("code", "")).casefold(), + str(error_fields.get("code", "")).casefold(), + str(error_fields.get("type", "")).casefold(), + str(error_fields.get("status", "")).casefold(), + } + if ( + codes + & {"provider_timeout", "timeout", "timeout_error", "deadline_exceeded"} + or "timeout" in lowered + ): self._set_terminal_hint("timeout") return if "cancel" in lowered: @@ -230,11 +242,7 @@ def _capture_terminal_hint(self, event_type: str, payload: dict[str, Any]) -> No "message_incomplete", "response.failed", } or isinstance(payload.get("error"), dict): - error = payload.get("error") - error_text = json.dumps(error).lower() - if "timeout" in error_text: - self._set_terminal_hint("timeout") - elif "cancel" in error_text: + if codes & {"cancelled", "canceled", "stream_cancelled"}: self._set_terminal_hint("cancelled") else: self._set_terminal_hint("provider_error") diff --git a/src/shim/gateway/streaming/session.py b/src/shim/gateway/streaming/session.py index d45ab03..f15561d 100644 --- a/src/shim/gateway/streaming/session.py +++ b/src/shim/gateway/streaming/session.py @@ -22,7 +22,7 @@ logger = logging.getLogger(__name__) -_DETACHED_FINALIZERS: set[asyncio.Task[Any]] = set() +_PROVIDER_CLOSE_TIMEOUT_SECONDS = 5.0 TerminalObserver = Callable[[StreamTerminalStatus], None] @@ -47,7 +47,11 @@ def __init__( clock: Callable[[], datetime] | None = None, parent_context: Context | None = None, terminal_observer: TerminalObserver | None = None, + finalization_tasks: set[asyncio.Task[Any]] | None = None, ) -> None: + self._finalization_tasks = ( + finalization_tasks if finalization_tasks is not None else set() + ) self.meter = meter self._finalizer = finalizer self._stream_start_recorder = stream_start_recorder @@ -109,8 +113,7 @@ async def aclose(self) -> None: finally: if self._terminal is None: self.meter.finish() - await self._close_provider_stream() - await self._finalize_safely("client_disconnected") + await self._cleanup("client_disconnected") async def record_stream_start(self) -> None: if self._stream_started: @@ -180,7 +183,12 @@ async def _iterate_stream(self) -> AsyncIterator[bytes]: finally: terminal = terminal or "client_disconnected" self.meter.finish() + await self._cleanup(terminal) + + async def _cleanup(self, terminal: StreamTerminalStatus) -> None: + try: await self._close_provider_stream() + finally: await self._finalize_safely(terminal) async def _rendered_chunks(self) -> AsyncIterator[bytes]: @@ -234,10 +242,11 @@ async def _finalize_safely( task = asyncio.create_task( self.finalize(terminal_status, error_message=error_message) ) + self._finalization_tasks.add(task) + task.add_done_callback(self._finalization_tasks.discard) try: await asyncio.shield(task) except asyncio.CancelledError: - _DETACHED_FINALIZERS.add(task) task.add_done_callback(_observe_detached_finalizer) logger.debug("Stream finalization continuing after response cancellation") raise @@ -250,19 +259,18 @@ async def _finalize_safely( async def _close_provider_stream(self) -> None: if self._provider_stream is None or self._provider_closed: return - self._provider_closed = True close = self._provider_close or getattr(self._provider_stream, "aclose", None) if close is None: close = getattr(self._provider_stream, "close", None) if not callable(close): return try: - result = close() - if isawaitable(result): - await result - except BaseException as exc: - if isinstance(exc, (KeyboardInterrupt, SystemExit)): - raise + async with asyncio.timeout(_PROVIDER_CLOSE_TIMEOUT_SECONDS): + result = close() + if isawaitable(result): + await result + self._provider_closed = True + except Exception as exc: logger.debug("Provider stream close failed type=%s", type(exc).__name__) def _terminal_from_hint(self) -> StreamTerminalStatus: @@ -307,7 +315,6 @@ def _terminal_error( def _observe_detached_finalizer(task: asyncio.Task[Any]) -> None: - _DETACHED_FINALIZERS.discard(task) if task.cancelled(): logger.error( "Detached stream finalization was cancelled; durable stale recovery retained" diff --git a/src/shim/gateway/usage.py b/src/shim/gateway/usage.py index f435493..816a2eb 100644 --- a/src/shim/gateway/usage.py +++ b/src/shim/gateway/usage.py @@ -2,15 +2,22 @@ from __future__ import annotations +import asyncio import json +import logging from datetime import datetime, timezone from decimal import Decimal -from threading import Lock +from queue import Full, Queue, ShutDown +from threading import Thread from typing import Literal, Protocol, TextIO, TypeAlias from shim.billing.pricing import DEFAULT_PRICE_BOOK, compute_cost_usd from shim.gateway.kernel.result import AdmissionState, PreparedInference from shim.gateway.streaming.finalization import StreamFinalization +from shim.observability.metrics import LOCAL_USAGE_DROPPED_TOTAL + + +logger = logging.getLogger(__name__) UsageFailureReason: TypeAlias = Literal[ @@ -61,11 +68,37 @@ async def fail( class LocalUsageLifecycle: - """Write one redacted JSONL event for each terminal local request.""" + """Queue redacted JSONL events; drop newest on overflow or sink failure.""" - def __init__(self, stream: TextIO) -> None: + def __init__(self, stream: TextIO, *, capacity: int = 1024) -> None: + if capacity < 1: + raise ValueError("event queue capacity must be positive") self._stream = stream - self._write_lock = Lock() + self._queue: Queue[str] = Queue(maxsize=capacity) + self._writer: Thread | None = None + self.dropped_events = 0 + self.write_failures = 0 + + async def aclose(self, timeout_seconds: float = 5.0) -> None: + self._queue.shutdown() + if self._writer is not None: + await asyncio.to_thread(self._writer.join, timeout_seconds) + + def _write_events(self) -> None: + while True: + try: + line = self._queue.get() + except ShutDown: + return + try: + self._stream.write(f"{line}\n") + self._stream.flush() + except Exception as exc: + self.write_failures += 1 + LOCAL_USAGE_DROPPED_TOTAL.labels(reason="sink_failure").inc() + logger.warning("Local usage event dropped type=%s", type(exc).__name__) + finally: + self._queue.task_done() async def admit( self, @@ -185,6 +218,13 @@ def _write( ), } line = json.dumps(event, ensure_ascii=False, separators=(",", ":")) - with self._write_lock: - self._stream.write(f"{line}\n") - self._stream.flush() + try: + self._queue.put_nowait(line) + except (Full, ShutDown): + # Drop newest telemetry; authoritative enterprise accounting is separate. + self.dropped_events += 1 + LOCAL_USAGE_DROPPED_TOTAL.labels(reason="queue_unavailable").inc() + return + if self._writer is None: + self._writer = Thread(target=self._write_events, daemon=True) + self._writer.start() diff --git a/src/shim/observability/metrics.py b/src/shim/observability/metrics.py index 6a313c2..2c85386 100644 --- a/src/shim/observability/metrics.py +++ b/src/shim/observability/metrics.py @@ -161,3 +161,10 @@ def _model_family(model: str) -> str: if folded.startswith(prefix): return f"{prefix}-*" return "other" + + +LOCAL_USAGE_DROPPED_TOTAL = Counter( + "shim_local_usage_dropped_total", + "Non-durable local usage events dropped by the bounded writer.", + ["reason"], +) diff --git a/src/shim/privacy/deanonymizer.py b/src/shim/privacy/deanonymizer.py index 8ca3417..a7ca324 100644 --- a/src/shim/privacy/deanonymizer.py +++ b/src/shim/privacy/deanonymizer.py @@ -226,11 +226,9 @@ def restore_events(self, payload: dict[str, Any]) -> list[dict[str, Any]]: ] def _restore_fragment(self, key: tuple[object, ...], fragment: str) -> str: - text = self._buffers.pop(key, "") + fragment - ready, carry = _split_placeholder_prefix(text, self._verification_map) - if carry: - self._buffers[key] = carry - return self._scrubber.deanonymize(ready, self._verification_map) + return restore_fragment( + self._buffers, key, fragment, self._verification_map, self._scrubber + ) def _flush_events(self, index: object | None = None) -> list[dict[str, Any]]: events: list[dict[str, Any]] = [] @@ -347,11 +345,9 @@ def restore_chat_chunk(self, payload: dict[str, Any]) -> dict[str, Any]: ) def _restore_fragment(self, key: tuple[object, ...], fragment: str) -> str: - text = self._buffers.pop(key, "") + fragment - ready, carry = _split_placeholder_prefix(text, self._verification_map) - if carry: - self._buffers[key] = carry - return self._scrubber.deanonymize(ready, self._verification_map) + return restore_fragment( + self._buffers, key, fragment, self._verification_map, self._scrubber + ) def _drop_buffers(self, event_prefix: str, item_id: object) -> None: for key in tuple(self._buffers): @@ -382,6 +378,20 @@ def _flush_chat_choice(self, choice_index: object, delta: dict[str, Any]) -> Non function["arguments"] = str(function.get("arguments") or "") + restored +def restore_fragment( + buffers: dict[tuple[object, ...], str], + key: tuple[object, ...], + fragment: str, + placeholders: Mapping[str, str], + scrubber: PIIScrubberService, +) -> str: + text = buffers.pop(key, "") + fragment + ready, carry = _split_placeholder_prefix(text, placeholders) + if carry: + buffers[key] = carry + return scrubber.deanonymize(ready, placeholders) + + def _split_placeholder_prefix( text: str, placeholders: Mapping[str, str], @@ -390,10 +400,15 @@ def _split_placeholder_prefix( if last_open < 0: return text, "" tail = text[last_open:] - compact_tail = "<" + "".join(tail[1:].split()) + identifier = tail[1:].strip() + compact_tail = "<" + identifier + if tail[-1:].isspace() and identifier and f"{compact_tail}>" not in placeholders: + return text, "" if any( len(compact_tail) < len(placeholder) and placeholder.startswith(compact_tail) for placeholder in placeholders ): + if len(tail) > 256: + raise ValueError("provider placeholder suffix exceeds 256 characters") return text[:last_open], tail return text, "" diff --git a/tests/gateway/pipeline/test_admission.py b/tests/gateway/pipeline/test_admission.py index 190d5f2..6e1fde7 100644 --- a/tests/gateway/pipeline/test_admission.py +++ b/tests/gateway/pipeline/test_admission.py @@ -247,3 +247,36 @@ def test_admission_rejects_invalid_injected_bounds(limits: dict[str, int]) -> No loop_detector=SimpleNamespace(check_exact_repeat=AsyncMock()), **values, ) + + +@pytest.mark.parametrize( + "provider,protocol,payload,expected", + [ + ("openai", "chat", {"n": 3, "generationConfig": {}}, 3), + ("openai", "responses", {"n": 3}, 1), + ("anthropic", "messages", {"n": 3}, 1), + ("google", "generate_content", {"n": 3}, 1), + ("google", "generate_content", {"generationConfig": {"candidateCount": 2}}, 2), + *[("openai", "chat", {"n": value}, 1) for value in (True, "3", None)], + ], +) +def test_native_candidate_counts(provider, protocol, payload, expected): + from shim.gateway.pipeline.admission import candidate_count + + assert ( + candidate_count( + SimpleNamespace(provider=provider, protocol=protocol, payload=payload) + ) + == expected + ) + + +@pytest.mark.parametrize("count", [0, -1, 10_001]) +def test_native_candidate_count_rejects_out_of_bounds(count): + from shim.gateway.pipeline.admission import candidate_count + + with pytest.raises(HTTPException) as error: + candidate_count( + SimpleNamespace(provider="openai", protocol="chat", payload={"n": count}) + ) + assert error.value.status_code == 400 diff --git a/tests/gateway/privacy/test_pii_scrubber.py b/tests/gateway/privacy/test_pii_scrubber.py index 2087783..1d985b6 100644 --- a/tests/gateway/privacy/test_pii_scrubber.py +++ b/tests/gateway/privacy/test_pii_scrubber.py @@ -1027,3 +1027,24 @@ def test_chat_stream_restores_split_tool_arguments( for chunk in (first, second) ) assert arguments == "alice@example.com" + + +def test_fragment_carry_is_bounded_and_preserves_every_placeholder_split(scrubber): + from shim.privacy.deanonymizer import restore_fragment + + placeholder, mapping = scrubber.scrub("alice@example.com") + for token in (placeholder, "< " + placeholder[1:-1] + " >"): + for split in range(len(token) + 1): + buffers = {} + output = restore_fragment(buffers, (0,), token[:split], mapping, scrubber) + output += restore_fragment(buffers, (0,), token[split:], mapping, scrubber) + assert output == "alice@example.com" + assert not buffers + buffers = {} + assert restore_fragment(buffers, (0,), "<", mapping, scrubber) == "" + with pytest.raises(ValueError, match="256"): + for _ in range(256): + restore_fragment(buffers, (0,), " ", mapping, scrubber) + assert len(buffers.get((0,), "")) <= 256 + for literal in ("1 < 2", "", " None: assert terminal.terminal_status == "completed" assert finalizer.await_count == 2 assert finalizer.await_args_list[0].args == finalizer.await_args_list[1].args + + +@pytest.mark.asyncio +@pytest.mark.parametrize("cancel_close", [False, True]) +async def test_blocked_provider_close_cannot_prevent_finalization( + monkeypatch, cancel_close +): + monkeypatch.setattr( + "shim.gateway.streaming.session._PROVIDER_CLOSE_TIMEOUT_SECONDS", 0.01 + ) + entered = asyncio.Event() + + async def close(): + entered.set() + await asyncio.Event().wait() + + finalizer = AsyncMock() + session = _session(finalizer) + session.bind(SimpleNamespace(), close=close) + task = asyncio.create_task(session.aclose()) + await entered.wait() + if cancel_close: + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + else: + await asyncio.wait_for(task, 1) + finalizer.assert_awaited_once() + assert not session._provider_closed + + +@pytest.mark.asyncio +async def test_shutdown_drains_cancelled_response_finalizer(): + from shim.gateway.pipeline.postprocess import ResponsePostprocessor + + entered = asyncio.Event() + release = asyncio.Event() + + async def finalize(terminal): + entered.set() + await release.wait() + + processor = ResponsePostprocessor( + SimpleNamespace(), heartbeat_interval_seconds=30, output_hash_salt=None + ) + session = StreamSession( + meter=StreamMeter( + provider="openai", requested_model="gpt-5.6-luna", prompt_tokens_estimated=1 + ), + finalizer=finalize, + stream_start_recorder=AsyncMock(), + finalization_tasks=processor._finalization_tasks, + ) + session.bind(SimpleNamespace(aclose=AsyncMock())) + task = asyncio.create_task(session.aclose()) + await entered.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + assert processor._finalization_tasks + release.set() + await processor.drain() + assert session.terminal_status == "client_disconnected" + assert not processor._finalization_tasks diff --git a/tests/gateway/test_admission_controls.py b/tests/gateway/test_admission_controls.py index 9bcb4a0..7f6757e 100644 --- a/tests/gateway/test_admission_controls.py +++ b/tests/gateway/test_admission_controls.py @@ -76,3 +76,23 @@ async def test_in_memory_admission_rejects_invalid_bounds() -> None: limit=2, window_seconds=1, ) + + +def test_expiry_index_handles_mixed_windows_and_bounds_evicted_entries(): + from shim.gateway.admission import _FixedWindowCounters + + now = 0.0 + counters = _FixedWindowCounters(max_entries=3, clock=lambda: now) + counters.increment("long", amount=1, window_seconds=100) + counters.increment("short", amount=1, window_seconds=1) + counters.increment("medium", amount=1, window_seconds=20) + now = 1.0 + counters.increment("new", amount=1, window_seconds=10) + assert list(counters.windows) == ["long", "medium", "new"] + for index in range(100): + counters.increment(index % 5, amount=1, window_seconds=10) + assert len(counters.windows) <= 3 + assert len(counters._expirations) <= 6 + now = 12.0 + assert counters.increment("fresh", amount=1, window_seconds=1) == 1 + assert list(counters.windows) == ["fresh"] diff --git a/tests/gateway/test_openai_sdk_transport.py b/tests/gateway/test_openai_sdk_transport.py index 19896f1..b10fbd5 100644 --- a/tests/gateway/test_openai_sdk_transport.py +++ b/tests/gateway/test_openai_sdk_transport.py @@ -1050,3 +1050,41 @@ def test_provider_boundary_defaults_match_sdk_limits() -> None: assert defaults[f"{provider}_READ_TIMEOUT_SECONDS"].default == 600 assert defaults[f"{provider}_WRITE_TIMEOUT_SECONDS"].default == 600 assert defaults[f"{provider}_POOL_TIMEOUT_SECONDS"].default == 600 + + +@pytest.mark.asyncio +async def test_chat_placeholder_overflow_on_finished_choice_is_a_terminal_error(): + async def chunks(): + yield SimpleNamespace( + model_dump=lambda **kwargs: { + "choices": [ + { + "index": 0, + "delta": {"content": "<" + " " * 256}, + "finish_reason": "stop", + } + ] + } + ) + + async with httpx.AsyncClient() as http: + execution = _execution(http, SimpleNamespace(save=AsyncMock())) + wire = b"".join( + [ + event + async for event in execution._chat_stream( + chunks(), + _prepared( + {"stream": True}, + protocol="chat", + tenant="11111111-1111-1111-1111-111111111111", + mapping={"": "alice@example.com"}, + ), + {"closed": False, "recorded": False}, + AsyncMock(), + ) + ] + ) + assert b"PROVIDER_UNAVAILABLE" in wire + assert b"[DONE]" not in wire + assert b"alice@example.com" not in wire diff --git a/tests/gateway/test_usage.py b/tests/gateway/test_usage.py index 23c3a07..8193bcf 100644 --- a/tests/gateway/test_usage.py +++ b/tests/gateway/test_usage.py @@ -66,6 +66,7 @@ async def test_local_usage_writes_one_exact_redacted_terminal_event() -> None: await lifecycle.heartbeat_stream(prepared) await lifecycle.finalize(prepared, _terminal()) + await lifecycle.aclose() lines = stream.getvalue().splitlines() assert len(lines) == 1 assert "secret-body" not in lines[0] @@ -108,11 +109,13 @@ async def test_local_usage_uses_null_cost_for_unsupported_model() -> None: stream = StringIO() prepared = _prepared(model="private-model") - await LocalUsageLifecycle(stream).finalize( + lifecycle = LocalUsageLifecycle(stream) + await lifecycle.finalize( prepared, _terminal(model="private-model"), ) + await lifecycle.aclose() assert json.loads(stream.getvalue())["estimated_cost_usd"] is None @@ -120,12 +123,99 @@ async def test_local_usage_uses_null_cost_for_unsupported_model() -> None: async def test_local_failure_writes_one_terminal_event() -> None: stream = StringIO() - await LocalUsageLifecycle(stream).fail( + lifecycle = LocalUsageLifecycle(stream) + await lifecycle.fail( _prepared(), reason="provider_rejected_without_usage", ) + await lifecycle.aclose() event = json.loads(stream.getvalue()) assert event["outcome"] == "provider_rejected_without_usage" assert event["completion_tokens"] == 0 assert len(stream.getvalue().splitlines()) == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", ["write", "flush"]) +async def test_local_sink_failure_does_not_fail_settlement(failure): + class BrokenSink(StringIO): + def write(self, value): + if failure == "write": + raise OSError("private sink detail") + return super().write(value) + + def flush(self): + if failure == "flush": + raise OSError("private sink detail") + + lifecycle = LocalUsageLifecycle(BrokenSink()) + await lifecycle.finalize(_prepared(), _terminal()) + await lifecycle.aclose() + assert lifecycle.write_failures == 1 + + +@pytest.mark.asyncio +async def test_blocked_sink_keeps_event_loop_and_buffer_bounded(): + import asyncio + from threading import Event + + entered = Event() + release = Event() + + class BlockedSink(StringIO): + def write(self, value): + entered.set() + release.wait(2) + return super().write(value) + + lifecycle = LocalUsageLifecycle(BlockedSink(), capacity=2) + try: + await lifecycle.finalize(_prepared(), _terminal()) + assert await asyncio.to_thread(entered.wait, 1) + for _ in range(10): + await lifecycle.finalize(_prepared(), _terminal()) + assert lifecycle._queue.qsize() == 2 + assert lifecycle.dropped_events == 8 + await asyncio.wait_for(asyncio.sleep(0), 0.1) + await lifecycle.aclose(timeout_seconds=0.01) + finally: + release.set() + await lifecycle.aclose() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("salt", [None, "hash-salt"]) +async def test_nonstream_hash_is_optional_without_changing_response(monkeypatch, salt): + from unittest.mock import AsyncMock, Mock + from fastapi.responses import JSONResponse + import shim.gateway.pipeline.postprocess as module + from shim.gateway.pipeline.provider_execution import ProviderNonStream + from shim.privacy.classification import content_ref + + payload = { + "nested": {"content": "İstanbul 🌍"}, + "usage": {"prompt_tokens": 2, "completion_tokens": 3}, + } + expected = JSONResponse(payload) + prepared = _prepared() + prepared.protocol = "chat" + prepared.admission.maximum_output_tokens = 10 + usage = SimpleNamespace(finalize=AsyncMock()) + dumps = Mock(wraps=json.dumps) + monkeypatch.setattr(json, "dumps", dumps) + response = await module.ResponsePostprocessor( + usage, heartbeat_interval_seconds=30, output_hash_salt=salt + ).finalize(prepared, ProviderNonStream(payload, "upstream-id"), stream_session=None) + assert response.body == expected.body + assert response.headers["x-request-id"] == "upstream-id" + usage.finalize.assert_awaited_once() + terminal = usage.finalize.await_args.args[1] + assert dumps.call_count == 1 + assert terminal.usage.output_hash == ( + content_ref( + salt, json.dumps(payload, ensure_ascii=False, separators=(",", ":")) + ) + if salt is not None + else None + )