"""Tests for inference domain models, error categories, and redaction. Proves that: - InferenceTarget serialization excludes raw auth secret values - InferenceResult serialization does not include raw auth headers - Sensitive headers in extra_headers are redacted when serialized for logging - Error messages don't leak bearer tokens or API keys - All error categories exist and have correct retryability defaults - Models serialize/deserialize correctly Requirements: 2.8, 2.9 """ from __future__ import annotations import uuid import pytest from services.shared.inference.errors import InferenceError, InferenceErrorCategory from services.shared.inference.models import ( ChatMessage, InferenceResult, InferenceTarget, ModelLineage, ProviderCapabilities, StructuredGenerationRequest, TokenUsage, ) from services.shared.inference.redaction import ( SENSITIVE_HEADER_NAMES, redact_error_message, redact_headers, redact_target_for_logging, ) # --------------------------------------------------------------------------- # Fixtures # --------------------------------------------------------------------------- def _make_target( *, auth_secret_ref: str | None = "vault://inference/openai-key", extra_headers: dict[str, str] | None = None, ) -> InferenceTarget: """Create a realistic InferenceTarget for testing.""" if extra_headers is None: extra_headers = { "Authorization": "Bearer sk-live-abc123xyz456", "X-Api-Key": "secret-key-99", "X-Request-Source": "stonks-oracle", } return InferenceTarget( endpoint_id=uuid.uuid4(), deployment_id=uuid.uuid4(), protocol="openai_chat", base_url="https://api.example.com/v1", model="gpt-4o-mini", capabilities=ProviderCapabilities( chat_completions=True, json_schema=True, usage=True, ), auth_secret_ref=auth_secret_ref, extra_headers=extra_headers, ) # --------------------------------------------------------------------------- # Task 10.3: Serialization tests — credentials excluded # --------------------------------------------------------------------------- class TestTargetRedaction: """InferenceTarget serialization never leaks secret values.""" def test_auth_secret_ref_shows_reference_name_only(self) -> None: """The auth_secret_ref field shows the reference path, not a resolved value.""" target = _make_target(auth_secret_ref="vault://inference/openai-key") serialized = redact_target_for_logging(target) # The ref name is preserved (so operators can identify which secret) assert serialized["auth_secret_ref"] == "vault://inference/openai-key" # But no raw secret value appears anywhere in the serialized output flat = str(serialized) assert "sk-live-abc123xyz456" not in flat def test_sensitive_headers_redacted_in_extra_headers(self) -> None: """Authorization, X-Api-Key, and other sensitive headers are redacted.""" target = _make_target( extra_headers={ "Authorization": "Bearer sk-live-abc123xyz456", "X-Api-Key": "secret-key-99", "X-Request-Source": "stonks-oracle", "Content-Type": "application/json", } ) serialized = redact_target_for_logging(target) headers = serialized["extra_headers"] # Sensitive headers show redacted placeholder assert headers["Authorization"] == "***REDACTED***" assert headers["X-Api-Key"] == "***REDACTED***" # Non-sensitive headers are preserved assert headers["X-Request-Source"] == "stonks-oracle" assert headers["Content-Type"] == "application/json" def test_no_raw_auth_in_full_serialized_output(self) -> None: """The full serialized dict does not contain any raw secret strings.""" target = _make_target( extra_headers={ "Authorization": "Bearer my-super-secret-token-12345678", "api-key": "ak_prod_9876543210abcdef", } ) serialized = redact_target_for_logging(target) flat = str(serialized).lower() assert "my-super-secret-token-12345678" not in flat assert "ak_prod_9876543210abcdef" not in flat def test_none_auth_secret_ref_serializes_as_none(self) -> None: """Targets without auth show None, not a placeholder.""" target = _make_target(auth_secret_ref=None, extra_headers={}) serialized = redact_target_for_logging(target) assert serialized["auth_secret_ref"] is None def test_case_insensitive_header_matching(self) -> None: """Header name matching is case-insensitive.""" target = _make_target( extra_headers={ "AUTHORIZATION": "Bearer token123", "x-API-KEY": "key456", } ) serialized = redact_target_for_logging(target) headers = serialized["extra_headers"] assert headers["AUTHORIZATION"] == "***REDACTED***" assert headers["x-API-KEY"] == "***REDACTED***" class TestInferenceResultSerialization: """InferenceResult model_dump does not include raw auth headers.""" def test_result_does_not_contain_auth_headers(self) -> None: """InferenceResult serialization has no field for raw auth data.""" result = InferenceResult( content='{"answer": 42}', parsed={"answer": 42}, endpoint_id=uuid.uuid4(), deployment_id=uuid.uuid4(), model="qwen-9b", protocol="openai_chat", structured_mode="json_schema", latency_ms=450, usage=TokenUsage(input_tokens=100, output_tokens=50, total_tokens=150), request_id="req-abc-123", retries=1, ) serialized = result.model_dump() flat = str(serialized).lower() # No auth/secret fields exist in the serialized output assert "auth" not in flat assert "secret" not in flat assert "bearer" not in flat assert "api_key" not in flat or "api-key" not in flat def test_result_includes_typed_metadata(self) -> None: """InferenceResult contains all required typed metadata fields.""" eid = uuid.uuid4() did = uuid.uuid4() result = InferenceResult( content="hello", endpoint_id=eid, deployment_id=did, model="test-model", protocol="ollama_native", structured_mode="prompt_only", latency_ms=200, usage=TokenUsage(input_tokens=10, output_tokens=5), request_id="req-xyz", repaired=True, retries=2, ) data = result.model_dump() assert data["endpoint_id"] == eid assert data["deployment_id"] == did assert data["model"] == "test-model" assert data["protocol"] == "ollama_native" assert data["structured_mode"] == "prompt_only" assert data["latency_ms"] == 200 assert data["usage"]["input_tokens"] == 10 assert data["usage"]["output_tokens"] == 5 assert data["request_id"] == "req-xyz" assert data["repaired"] is True assert data["retries"] == 2 class TestErrorMessageRedaction: """Error messages don't leak bearer tokens or API keys.""" def test_bearer_token_redacted(self) -> None: """Bearer tokens are replaced in error messages.""" msg = "Authentication failed with Bearer sk-live-abc123xyz456def789" redacted = redact_error_message(msg) assert "sk-live-abc123xyz456def789" not in redacted assert "***REDACTED***" in redacted def test_api_key_prefix_redacted(self) -> None: """Strings matching API key patterns are redacted.""" msg = "Invalid key: api_key_abcdef1234567890abcdef" redacted = redact_error_message(msg) assert "api_key_abcdef1234567890abcdef" not in redacted assert "***REDACTED***" in redacted def test_long_secret_like_string_redacted(self) -> None: """Strings matching API key prefix patterns are redacted.""" secret = "sk-prod_abcdef1234567890abcdef1234567890xyz" msg = f"Connection refused for endpoint token={secret}" redacted = redact_error_message(msg) assert secret not in redacted def test_short_strings_preserved(self) -> None: """Short normal words are not false-positive redacted.""" msg = "Connection timeout after 30 seconds to endpoint" redacted = redact_error_message(msg) assert redacted == msg def test_multiple_secrets_all_redacted(self) -> None: """Multiple secrets in one message are all replaced.""" msg = "Bearer sk-test-aabbccddee123456 failed, also key-prod_xyzxyzxyzxyz1234" redacted = redact_error_message(msg) assert "sk-test-aabbccddee123456" not in redacted assert "key-prod_xyzxyzxyzxyz1234" not in redacted # --------------------------------------------------------------------------- # Error category tests # --------------------------------------------------------------------------- class TestErrorCategories: """All expected error categories exist with correct defaults.""" def test_required_categories_present(self) -> None: """All spec-required categories are covered by the enum.""" # The enum may contain additional granular categories beyond the spec # requirement, but must cover: timeout, authentication, rate limit, # server, invalid response, schema, capability, policy, connection, unknown actual_values = {c.value for c in InferenceErrorCategory} # Required base categories (may be named differently for granularity) assert "timeout" in actual_values assert "server_error" in actual_values assert "invalid_response" in actual_values assert "policy_violation" in actual_values assert "connection_error" in actual_values assert "unknown" in actual_values # Auth-related assert any("auth" in v for v in actual_values) # Rate limit assert any("rate" in v for v in actual_values) # Schema/validation assert any("schema" in v or "violation" in v for v in actual_values) # Capability assert any("capability" in v or "unavailable" in v for v in actual_values) def test_retryable_categories(self) -> None: """Timeout, rate_limit, server_error, connection_error default to retryable.""" err_timeout = InferenceError(InferenceErrorCategory.TIMEOUT, "timed out") err_rate = InferenceError(InferenceErrorCategory.RATE_LIMITED, "429") err_server = InferenceError(InferenceErrorCategory.SERVER_ERROR, "500") err_conn = InferenceError(InferenceErrorCategory.CONNECTION_ERROR, "refused") assert err_timeout.retryable is True assert err_rate.retryable is True assert err_server.retryable is True assert err_conn.retryable is True def test_non_retryable_categories(self) -> None: """Auth, schema, capability, policy, invalid_response, unknown default non-retryable.""" err_auth = InferenceError(InferenceErrorCategory.AUTH_FAILED, "401") err_schema = InferenceError(InferenceErrorCategory.SCHEMA_VIOLATION, "bad") err_cap = InferenceError( InferenceErrorCategory.CAPABILITY_UNAVAILABLE, "no json" ) err_policy = InferenceError(InferenceErrorCategory.POLICY_VIOLATION, "denied") err_invalid = InferenceError( InferenceErrorCategory.INVALID_RESPONSE, "malformed" ) err_unknown = InferenceError(InferenceErrorCategory.UNKNOWN, "???") assert err_auth.retryable is False assert err_schema.retryable is False assert err_cap.retryable is False assert err_policy.retryable is False assert err_invalid.retryable is False assert err_unknown.retryable is False def test_retryable_property_on_category(self) -> None: """The retryable property is accessible directly on the category enum.""" assert InferenceErrorCategory.TIMEOUT.retryable is True assert InferenceErrorCategory.AUTH_FAILED.retryable is False def test_error_includes_status_code(self) -> None: """HTTP status code is preserved on the error.""" err = InferenceError( InferenceErrorCategory.RATE_LIMITED, "Too many requests", status_code=429, ) assert err.status_code == 429 def test_error_str_format(self) -> None: """String representation includes the message.""" err = InferenceError(InferenceErrorCategory.TIMEOUT, "Request timed out") assert "Request timed out" in str(err) # --------------------------------------------------------------------------- # Header redaction utility # --------------------------------------------------------------------------- class TestRedactHeaders: """Direct header redaction function tests.""" def test_all_sensitive_names_redacted(self) -> None: """Every name in SENSITIVE_HEADER_NAMES is redacted.""" headers = {name: f"value-for-{name}" for name in SENSITIVE_HEADER_NAMES} redacted = redact_headers(headers) for name in SENSITIVE_HEADER_NAMES: assert redacted[name] == "***REDACTED***" def test_non_sensitive_preserved(self) -> None: """Non-sensitive headers pass through unchanged.""" headers = {"Content-Type": "application/json", "Accept": "text/html"} redacted = redact_headers(headers) assert redacted == headers def test_empty_headers(self) -> None: """Empty dict returns empty dict.""" assert redact_headers({}) == {} # --------------------------------------------------------------------------- # Model type tests # --------------------------------------------------------------------------- class TestModelTypes: """Basic validation of model types.""" def test_chat_message_serialization(self) -> None: """ChatMessage serializes correctly.""" msg = ChatMessage(role="user", content="Hello") data = msg.model_dump() assert data == {"role": "user", "content": "Hello"} def test_token_usage_defaults(self) -> None: """TokenUsage fields default to None.""" usage = TokenUsage() assert usage.input_tokens is None assert usage.output_tokens is None assert usage.total_tokens is None def test_structured_request_defaults(self) -> None: """StructuredGenerationRequest has sensible defaults.""" req = StructuredGenerationRequest( messages=[ChatMessage(role="user", content="test")], max_output_tokens=512, ) assert req.temperature == 0.0 assert req.seed == 0 assert req.timeout_seconds == 120.0 assert req.trace_id == "" assert req.json_schema is None def test_provider_capabilities_immutable(self) -> None: """ProviderCapabilities is frozen.""" caps = ProviderCapabilities(chat_completions=True) with pytest.raises(Exception): caps.chat_completions = False # type: ignore[misc] def test_inference_target_immutable(self) -> None: """InferenceTarget is frozen.""" target = _make_target() with pytest.raises(Exception): target.model = "other" # type: ignore[misc] def test_model_lineage_serialization(self) -> None: """ModelLineage serializes all fields.""" lineage = ModelLineage( endpoint_id=uuid.uuid4(), deployment_id=uuid.uuid4(), model="qwen-9b", protocol="openai_chat", structured_mode="json_schema", request_id="req-123", latency_ms=300, retries=0, trace_id="trace-abc", ) data = lineage.model_dump() assert data["model"] == "qwen-9b" assert data["trace_id"] == "trace-abc"