diff --git a/app.py b/app.py index ac90ed789..bae972452 100644 --- a/app.py +++ b/app.py @@ -83,14 +83,23 @@ from starlette.responses import RedirectResponse # ========= LOGGING ========= import logging.handlers from core.constants import DATA_DIR +from core.log_safety import ( + CAPABILITY_DIAGNOSTICS_LOGGER, + ScopedDiagnosticsFilter, + application_log_settings, +) _root_logger = logging.getLogger() _log_level_name = os.getenv("LOG_LEVEL", "INFO").strip().upper() -_log_level = getattr(logging, _log_level_name, logging.INFO) -if not isinstance(_log_level, int): - _log_level_name = "INFO" - _log_level = logging.INFO -_root_logger.setLevel(_log_level) +_application_log_level, _capability_debug = application_log_settings(_log_level_name) +_root_logger.setLevel(_application_log_level) +logging.getLogger(CAPABILITY_DIAGNOSTICS_LOGGER).setLevel( + logging.DEBUG if _capability_debug else logging.NOTSET +) +_diagnostics_filter = ScopedDiagnosticsFilter( + _application_log_level, + capability_debug=_capability_debug, +) _formatter = logging.Formatter('%(asctime)s - %(name)s - %(levelname)s - %(message)s') # Clear existing handlers to avoid duplicates @@ -98,7 +107,8 @@ for _h in list(_root_logger.handlers): _root_logger.removeHandler(_h) _console_h = logging.StreamHandler() -_console_h.setLevel(_log_level) +_console_h.setLevel(logging.DEBUG) +_console_h.addFilter(_diagnostics_filter) _console_h.setFormatter(_formatter) _root_logger.addHandler(_console_h) @@ -113,7 +123,8 @@ try: _file_h = logging.handlers.RotatingFileHandler( _log_file, maxBytes=5 * 1024 * 1024, backupCount=3, encoding="utf-8" ) - _file_h.setLevel(_log_level) + _file_h.setLevel(logging.DEBUG) + _file_h.addFilter(_diagnostics_filter) _file_h.setFormatter(_formatter) _root_logger.addHandler(_file_h) except Exception as e: @@ -1285,4 +1296,4 @@ if __name__ == "__main__": bind_host = os.getenv("APP_BIND", "127.0.0.1") bind_port = int(os.getenv("APP_PORT", "7000")) - uvicorn.run(app, host=bind_host, port=bind_port, log_level=_log_level_name.lower()) + uvicorn.run(app, host=bind_host, port=bind_port, log_level=_application_log_level) diff --git a/core/log_safety.py b/core/log_safety.py index 2339a73b6..a314e1e98 100644 --- a/core/log_safety.py +++ b/core/log_safety.py @@ -7,9 +7,63 @@ raw leaks those secrets, so route/diagnostic logs run URLs through also doubles as a sanitizer barrier for CodeQL's clear-text-logging query. """ +from __future__ import annotations + +import logging + from urllib.parse import urlparse, urlunparse +CAPABILITY_DIAGNOSTICS_LOGGER = "src.model_capability_readers" + +_LOG_LEVELS = { + "DEBUG": logging.DEBUG, + "INFO": logging.INFO, + "WARN": logging.WARNING, + "WARNING": logging.WARNING, + "ERROR": logging.ERROR, + "FATAL": logging.CRITICAL, + "CRITICAL": logging.CRITICAL, +} + + +def application_log_settings(value: object) -> tuple[int, bool]: + """Return the safe app level and whether scoped capability debug is on. + + Application-wide DEBUG logging can expose request bodies, provider + responses, or credentials from unrelated libraries. The model capability + catalog has a deliberately bounded DEBUG summary, so a DEBUG request is + translated into INFO for the application and enabled only for that logger. + Unknown values also fail closed to INFO. + """ + + requested = _LOG_LEVELS.get(str(value or "INFO").strip().upper(), logging.INFO) + return max(requested, logging.INFO), requested == logging.DEBUG + + +class ScopedDiagnosticsFilter(logging.Filter): + """Allow normal application records plus one explicitly scoped DEBUG log.""" + + def __init__( + self, + application_level: int, + *, + capability_debug: bool = False, + ) -> None: + super().__init__() + self.application_level = application_level + self.capability_debug = capability_debug + + def filter(self, record: logging.LogRecord) -> bool: + if record.levelno >= self.application_level: + return True + return ( + self.capability_debug + and record.levelno >= logging.DEBUG + and record.name == CAPABILITY_DIAGNOSTICS_LOGGER + ) + + def redact_url(url: str) -> str: """Return a URL safe for logs by removing userinfo and query/fragment. diff --git a/src/model_capability_readers/__init__.py b/src/model_capability_readers/__init__.py index 5a67edf3f..d25721d27 100644 --- a/src/model_capability_readers/__init__.py +++ b/src/model_capability_readers/__init__.py @@ -108,25 +108,71 @@ def records_from_payload( if vendor_id == pcs.PROVIDER_UNKNOWN: vendor_id = detect_vendor(base_url, endpoint_kind) reader = reader_for_vendor(vendor_id) - if reader is generic_openai: - record_vendor = vendor_id if vendor_id else VENDOR_UNKNOWN - records = reader.records_from_payload( + + record_vendor = vendor_id if vendor_id else VENDOR_UNKNOWN + + def annotate( + records: tuple[ModelCapabilityRecord, ...], + *, + shape_id: str, + fallback: bool, + ) -> tuple[ModelCapabilityRecord, ...]: + return tuple( + replace( + record, + provider_source=resolution.provider_source, + catalog_shape_id=shape_id, + fallback=fallback, + ) + for record in records + ) + + shape = pcs.catalog_shape_for_id(resolution.shape_id) + if resolution.fallback or reader is generic_openai or shape is None: + records = generic_openai.records_from_payload( payload, vendor_id=record_vendor, endpoint_id=endpoint_id, base_url=base_url, ) - else: - records = reader.records_from_payload(payload, endpoint_id=endpoint_id, base_url=base_url) - normalized = tuple( - replace( - record, - provider_source=resolution.provider_source, - catalog_shape_id=resolution.shape_id, + normalized = annotate( + records, + shape_id=resolution.shape_id, fallback=resolution.fallback, ) - for record in records - ) + else: + normalized_records: list[ModelCapabilityRecord] = [] + for item in shape.items(payload): + item_payload = shape.payload_for_item(payload, item) + if shape.item_matches(item): + native_records = reader.records_from_payload( + item_payload, + endpoint_id=endpoint_id, + base_url=base_url, + ) + if native_records: + normalized_records.extend( + annotate(native_records, shape_id=shape.shape_id, fallback=False) + ) + continue + + fallback_record = generic_openai.record_from_model( + item, + vendor_id=record_vendor, + endpoint_id=endpoint_id, + base_url=base_url, + ) + if fallback_record: + fallback_shape = pcs.fallback_shape_for_payload(item_payload) + normalized_records.extend( + annotate( + (fallback_record,), + shape_id=fallback_shape.shape_id if fallback_shape else "", + fallback=True, + ) + ) + normalized = tuple(normalized_records) + if logger.isEnabledFor(logging.DEBUG): families = sorted({record.capability.family for record in normalized}) features = sorted( @@ -144,16 +190,27 @@ def records_from_payload( if control.control } ) + fallback_count = sum(record.fallback for record in normalized) + diagnostic_provider = ( + resolution.provider_id + if resolution.provider_id in pcs.PROVIDER_SCHEMAS + else "unknown" + if resolution.provider_id == pcs.PROVIDER_UNKNOWN + else "unregistered" + ) logger.debug( "[model-capability] normalized: canonical_version=%s provider=%s " "provider_source=%s catalog_shape=%s fallback=%s records=%d " + "native_records=%d fallback_records=%d " "families=%s features=%s controls=%s", CANONICAL_MODEL_SHAPE_VERSION, - resolution.provider_id, + diagnostic_provider, resolution.provider_source, resolution.shape_id or "unknown", - resolution.fallback, + bool(fallback_count), len(normalized), + len(normalized) - fallback_count, + fallback_count, families, features, controls, diff --git a/src/model_capability_readers/sglang.py b/src/model_capability_readers/sglang.py index 3ac59b060..be1ce69af 100644 --- a/src/model_capability_readers/sglang.py +++ b/src/model_capability_readers/sglang.py @@ -3,6 +3,7 @@ from __future__ import annotations from collections.abc import Mapping +from dataclasses import replace from pathlib import PurePosixPath from typing import Any @@ -15,6 +16,7 @@ from src.model_capability_readers.base import ( build_capability, compact_str, deterministic_controls_from_supported_parameters, + int_limit, openai_model_items, stable_model_id_for, ) @@ -95,5 +97,14 @@ def records_from_payload( base_url=base_url, ) if record: + context_tokens = int_limit(item.get("max_model_len")) + if context_tokens: + record = replace( + record, + capability=build_capability( + family=mc.FAMILY_UNKNOWN, + limits={"context_tokens": context_tokens}, + ), + ) records.append(record) return tuple(records) diff --git a/src/provider_capability_schemas.py b/src/provider_capability_schemas.py index 55f4809da..bafc06c05 100644 --- a/src/provider_capability_schemas.py +++ b/src/provider_capability_schemas.py @@ -83,36 +83,49 @@ class ProviderCatalogShape: def items(self, payload: Any) -> tuple[Mapping[str, Any], ...]: return _items_for_envelope(payload, self.envelope) + def item_matches(self, item: Mapping[str, Any]) -> bool: + if self.identity_paths and not any( + (value := _path_value(item, path)) is not _MISSING + and value is not None + and value != "" + for path in self.identity_paths + ): + return False + if not all(_path_present(item, path) for path in self.required_item_paths): + return False + if self.required_item_any_paths and not any( + _path_present(item, path) for path in self.required_item_any_paths + ): + return False + if any( + not isinstance(_path_value(item, path), expected_types) + for path, expected_types in self.item_types + ): + return False + if any(_path_value(item, path) not in expected for path, expected in self.item_values): + return False + return True + def matches(self, payload: Any) -> bool: if self.required_root_paths: if not isinstance(payload, Mapping): return False if not all(_path_present(payload, path) for path in self.required_root_paths): return False + return any(self.item_matches(item) for item in self.items(payload)) - for item in self.items(payload): - if self.identity_paths and not any( - (value := _path_value(item, path)) is not _MISSING - and value is not None - and value != "" - for path in self.identity_paths - ): - continue - if not all(_path_present(item, path) for path in self.required_item_paths): - continue - if self.required_item_any_paths and not any( - _path_present(item, path) for path in self.required_item_any_paths - ): - continue - if any( - not isinstance(_path_value(item, path), expected_types) - for path, expected_types in self.item_types - ): - continue - if any(_path_value(item, path) not in expected for path, expected in self.item_values): - continue - return True - return False + def payload_for_item(self, payload: Any, item: Mapping[str, Any]) -> Any: + """Return a one-item payload in the same provider-native envelope.""" + + if self.envelope == ENVELOPE_BARE_LIST: + return [item] + if self.envelope == ENVELOPE_SINGLE: + return item + if isinstance(payload, Mapping): + narrowed = dict(payload) + narrowed[self.envelope] = [item] + return narrowed + return {self.envelope: [item]} @dataclass(frozen=True) @@ -183,11 +196,19 @@ OPENROUTER_MODELS_SHAPE = ProviderCatalogShape( provider_id="openrouter", envelope=ENVELOPE_DATA, identity_paths=("id",), - required_item_any_paths=( + required_item_paths=( "architecture", + "canonical_slug", + "pricing", "supported_parameters", "top_provider", - "canonical_slug", + ), + item_types=( + ("architecture", (Mapping,)), + ("canonical_slug", (str,)), + ("pricing", (Mapping,)), + ("supported_parameters", (list, tuple)), + ("top_provider", (Mapping,)), ), detection_priority=90, ) @@ -232,7 +253,7 @@ LMSTUDIO_MODELS_V1_SHAPE = ProviderCatalogShape( "quantization", ), item_types=(("type", (str,)),), - detection_priority=100, + detection_priority=0, ) LMSTUDIO_MODELS_V0_SHAPE = ProviderCatalogShape( shape_id="lmstudio.models.native.v0", @@ -242,7 +263,7 @@ LMSTUDIO_MODELS_V0_SHAPE = ProviderCatalogShape( required_item_paths=("type",), required_item_any_paths=("arch", "compatibility_type", "state", "max_context_length"), item_types=(("type", (str,)),), - detection_priority=80, + detection_priority=0, ) LLAMACPP_PROPS_SHAPE = ProviderCatalogShape( shape_id="llamacpp.props.v1", @@ -268,7 +289,7 @@ MISTRAL_MODELS_SHAPE = ProviderCatalogShape( "capabilities.classification", ), item_types=(("capabilities", (Mapping,)),), - detection_priority=100, + detection_priority=0, ) COPILOT_MODELS_SHAPE = ProviderCatalogShape( shape_id="github-copilot.models.v1", @@ -353,7 +374,7 @@ COHERE_MODELS_SHAPE = ProviderCatalogShape( "sampling_defaults", ), item_types=(("endpoints", (list, tuple)),), - detection_priority=100, + detection_priority=0, ) MINIMAX_MODELS_SHAPE = ProviderCatalogShape( shape_id="minimax.models.identity.v1", @@ -527,6 +548,18 @@ def schema_for_provider(value: Any) -> ProviderCapabilitySchema: return PROVIDER_SCHEMAS.get(normalize_provider_id(value), UNKNOWN_SCHEMA) +def provider_from_endpoint_kind(value: Any) -> str: + """Return a provider only for registered provider-valued endpoint kinds. + + Endpoint configuration normally stores transport categories such as + ``auto``, ``local``, ``api``, and ``proxy``. Those categories and unknown + values must not preempt provider identity from a host or native payload. + """ + + provider_id = normalize_provider_id(value) + return provider_id if provider_id in PROVIDER_SCHEMAS else PROVIDER_UNKNOWN + + def _host_matches(host: str, suffix: str) -> bool: return host == suffix or host.endswith("." + suffix) @@ -576,6 +609,18 @@ def native_shape_for_payload( return sorted(best, key=lambda shape: shape.shape_id)[0] +def catalog_shape_for_id(shape_id: Any) -> ProviderCatalogShape | None: + return next( + ( + shape + for schema in PROVIDER_SCHEMAS.values() + for shape in schema.catalog_shapes + if shape.shape_id == shape_id + ), + None, + ) + + def fallback_shape_for_payload(payload: Any) -> ProviderCatalogShape | None: return next((shape for shape in FALLBACK_CATALOG_SHAPES if shape.matches(payload)), None) @@ -591,7 +636,7 @@ def resolve_provider( provider_source = PROVIDER_SOURCE_EXPLICIT if provider_id == PROVIDER_UNKNOWN: - provider_id = normalize_provider_id(endpoint_kind) + provider_id = provider_from_endpoint_kind(endpoint_kind) provider_source = PROVIDER_SOURCE_ENDPOINT_KIND if provider_id == PROVIDER_UNKNOWN: provider_id = provider_from_host(base_url) @@ -647,9 +692,11 @@ __all__ = [ "ProviderCapabilitySchema", "ProviderCatalogShape", "ProviderResolution", + "catalog_shape_for_id", "fallback_shape_for_payload", "native_shape_for_payload", "normalize_provider_id", + "provider_from_endpoint_kind", "provider_from_host", "resolve_provider", "schema_for_provider", diff --git a/tests/test_log_safety.py b/tests/test_log_safety.py index b806516d2..a4de26666 100644 --- a/tests/test_log_safety.py +++ b/tests/test_log_safety.py @@ -1,4 +1,13 @@ -from core.log_safety import redact_url +import logging + +import pytest + +from core.log_safety import ( + CAPABILITY_DIAGNOSTICS_LOGGER, + ScopedDiagnosticsFilter, + application_log_settings, + redact_url, +) def test_strips_userinfo(): @@ -30,3 +39,48 @@ def test_empty_and_none(): def test_garbage_does_not_raise(): # urlparse is lenient; just assert no credential-looking userinfo survives. assert "@" not in redact_url("::::not a url::::") + + +@pytest.mark.parametrize( + ("configured", "expected_level", "expected_capability_debug"), + ( + ("DEBUG", logging.INFO, True), + ("debug", logging.INFO, True), + ("INFO", logging.INFO, False), + ("WARNING", logging.WARNING, False), + ("ERROR", logging.ERROR, False), + ("CRITICAL", logging.CRITICAL, False), + ("not-a-level", logging.INFO, False), + (None, logging.INFO, False), + ), +) +def test_application_log_settings_scope_debug_and_fail_closed( + configured, + expected_level, + expected_capability_debug, +): + assert application_log_settings(configured) == ( + expected_level, + expected_capability_debug, + ) + + +def _record(name: str, level: int) -> logging.LogRecord: + return logging.LogRecord(name, level, __file__, 1, "message", (), None) + + +def test_scoped_diagnostics_filter_allows_only_bounded_debug_logger(): + log_filter = ScopedDiagnosticsFilter(logging.INFO, capability_debug=True) + + assert log_filter.filter(_record(CAPABILITY_DIAGNOSTICS_LOGGER, logging.DEBUG)) + assert log_filter.filter(_record("unrelated.library", logging.INFO)) + assert not log_filter.filter(_record("unrelated.library", logging.DEBUG)) + assert not log_filter.filter(_record(f"{CAPABILITY_DIAGNOSTICS_LOGGER}.raw", logging.DEBUG)) + + +def test_scoped_diagnostics_filter_respects_higher_application_level(): + log_filter = ScopedDiagnosticsFilter(logging.WARNING, capability_debug=False) + + assert log_filter.filter(_record("application", logging.WARNING)) + assert not log_filter.filter(_record("application", logging.INFO)) + assert not log_filter.filter(_record(CAPABILITY_DIAGNOSTICS_LOGGER, logging.DEBUG)) diff --git a/tests/test_model_capability_diagnostics.py b/tests/test_model_capability_diagnostics.py index 3582ae425..7b1969430 100644 --- a/tests/test_model_capability_diagnostics.py +++ b/tests/test_model_capability_diagnostics.py @@ -10,7 +10,10 @@ def test_normalization_debug_log_reports_shape_without_payload_identity(caplog): { "id": "sensitive-model-id", "architecture": {"modality": "text+image->text"}, + "canonical_slug": "provider/sensitive-model-id", + "pricing": {"prompt": "0.1", "completion": "0.2"}, "supported_parameters": ["tools", "temperature"], + "top_provider": {"context_length": 32768}, "private_field": "secret-value", } ] @@ -28,6 +31,8 @@ def test_normalization_debug_log_reports_shape_without_payload_identity(caplog): assert "catalog_shape=openrouter.models.rich.v1" in message assert "fallback=False" in message assert "records=1" in message + assert "native_records=1" in message + assert "fallback_records=0" in message assert "families=['chat']" in message assert "features=['tool_call', 'vision']" in message assert "controls=['temperature']" in message @@ -43,9 +48,11 @@ def test_fallback_debug_log_is_explicit_and_has_no_capability_claims(caplog): assert records[0].capability.capabilities == () message = caplog.messages[-1] - assert "provider=future_provider" in message + assert "provider=unregistered" in message assert "catalog_shape=fallback.models.list.v1" in message assert "fallback=True" in message + assert "native_records=0" in message + assert "fallback_records=1" in message assert "features=[]" in message @@ -53,7 +60,8 @@ def test_web_app_logging_uses_existing_log_level_environment_toggle(): source = (Path(__file__).resolve().parents[1] / "app.py").read_text(encoding="utf-8") assert 'os.getenv("LOG_LEVEL", "INFO")' in source - assert "_root_logger.setLevel(_log_level)" in source - assert "_console_h.setLevel(_log_level)" in source - assert "_file_h.setLevel(_log_level)" in source - assert "log_level=_log_level_name.lower()" in source + assert "application_log_settings(_log_level_name)" in source + assert "_root_logger.setLevel(_application_log_level)" in source + assert "_console_h.addFilter(_diagnostics_filter)" in source + assert "_file_h.addFilter(_diagnostics_filter)" in source + assert "log_level=_application_log_level" in source diff --git a/tests/test_provider_capability_schemas.py b/tests/test_provider_capability_schemas.py index 9d5349c85..87d4244a7 100644 --- a/tests/test_provider_capability_schemas.py +++ b/tests/test_provider_capability_schemas.py @@ -14,6 +14,22 @@ from src.model_capability_readers import ( ) +def _openrouter_payload(*items): + return { + "data": list(items) + or [ + { + "id": "provider/model", + "architecture": {"modality": "text->text"}, + "canonical_slug": "provider/model", + "pricing": {"prompt": "0.1", "completion": "0.2"}, + "supported_parameters": ["tools", "temperature"], + "top_provider": {"context_length": 32768}, + } + ] + } + + def test_provider_identity_and_catalog_shape_are_resolved_separately(): google_payload = { "models": [ @@ -71,6 +87,27 @@ def test_provider_host_matching_rejects_lookalikes_and_does_not_use_ports(): assert pcs.provider_from_host("http://127.0.0.1:30000") == pcs.PROVIDER_UNKNOWN +def test_transport_endpoint_kinds_do_not_preempt_host_or_payload_provider_identity(): + payload = _openrouter_payload() + + for endpoint_kind in ("auto", "local", "api", "proxy", "future-transport"): + from_host = pcs.resolve_provider( + payload, + endpoint_kind=endpoint_kind, + base_url="https://api.openrouter.ai/v1", + ) + from_payload = pcs.resolve_provider(payload, endpoint_kind=endpoint_kind) + + assert from_host.provider_id == "openrouter" + assert from_host.provider_source == pcs.PROVIDER_SOURCE_HOST + assert from_payload.provider_id == "openrouter" + assert from_payload.provider_source == pcs.PROVIDER_SOURCE_PAYLOAD + + assert pcs.provider_from_endpoint_kind("ollama") == "ollama" + assert pcs.provider_from_endpoint_kind("llama.cpp") == "llamacpp" + assert pcs.provider_from_endpoint_kind("proxy") == pcs.PROVIDER_UNKNOWN + + def test_provider_aliases_only_normalize_explicit_identity(): assert pcs.normalize_provider_id("opencode-go") == "opencode" assert pcs.normalize_provider_id("opencode-zen") == "opencode" @@ -102,8 +139,13 @@ def test_unregistered_explicit_provider_is_preserved_but_stays_on_fallback(): assert records[0].capability.capabilities == () -def test_current_native_catalog_shapes_are_discriminating(): +def test_native_catalog_shapes_resolve_with_required_provider_context(): cases = ( + ( + _openrouter_payload(), + "openrouter", + "openrouter.models.rich.v1", + ), ( {"models": [{"key": "local/model", "type": "llm", "capabilities": {"vision": True}}]}, "lmstudio", @@ -198,30 +240,135 @@ def test_current_native_catalog_shapes_are_discriminating(): ), ) + explicit_context_providers = {"cohere", "lmstudio", "mistral"} for payload, expected_provider, expected_shape in cases: - resolution = pcs.resolve_provider(payload) + explicit_provider = ( + expected_provider if expected_provider in explicit_context_providers else None + ) + resolution = pcs.resolve_provider(payload, provider=explicit_provider) assert resolution.provider_id == expected_provider - assert resolution.provider_source == pcs.PROVIDER_SOURCE_PAYLOAD + assert resolution.provider_source == ( + pcs.PROVIDER_SOURCE_EXPLICIT + if explicit_provider + else pcs.PROVIDER_SOURCE_PAYLOAD + ) assert resolution.shape_id == expected_shape assert resolution.fallback is False -def test_wrong_native_field_types_degrade_to_explicit_fallback_inventory(): - malformed_cohere = pcs.resolve_provider( - {"models": [{"name": "future", "endpoints": "chat", "context_length": 4096}]} - ) - malformed_mistral = pcs.resolve_provider( - {"data": [{"id": "future", "capabilities": ["completion_chat"]}]} +def test_ambiguous_common_fields_do_not_infer_provider_from_payload_alone(): + cases = ( + ({"data": [{"id": "generic", "architecture": {}}]}, "openrouter"), + ({"data": [{"id": "generic", "supported_parameters": ["tools"]}]}, "openrouter"), + ( + {"data": [{"id": "generic", "capabilities": {"completion_chat": True}}]}, + "mistral", + ), + ({"data": [{"id": "generic", "type": "llm", "arch": "future"}]}, "lmstudio"), + ( + {"models": [{"key": "generic", "type": "llm", "capabilities": {}}]}, + "lmstudio", + ), + ( + {"models": [{"name": "generic", "endpoints": ["chat"], "context_length": 4096}]}, + "cohere", + ), ) - assert malformed_cohere.provider_id == pcs.PROVIDER_UNKNOWN + for payload, provider_id in cases: + inferred = pcs.resolve_provider(payload) + contextual = pcs.resolve_provider(payload, provider=provider_id) + + assert inferred.provider_id == pcs.PROVIDER_UNKNOWN + assert inferred.fallback is True + assert contextual.provider_id == provider_id + + +def test_openrouter_payload_detection_requires_the_compound_official_shape(): + complete = _openrouter_payload() + assert pcs.resolve_provider(complete).shape_id == "openrouter.models.rich.v1" + + item = complete["data"][0] + for required_field in ( + "architecture", + "canonical_slug", + "pricing", + "supported_parameters", + "top_provider", + ): + partial = _openrouter_payload( + {key: value for key, value in item.items() if key != required_field} + ) + resolution = pcs.resolve_provider(partial) + + assert resolution.provider_id == pcs.PROVIDER_UNKNOWN + assert resolution.shape_id == "fallback.models.data.v1" + assert resolution.fallback is True + + +def test_wrong_native_field_types_degrade_to_explicit_fallback_inventory(): + malformed_cohere = pcs.resolve_provider( + {"models": [{"name": "future", "endpoints": "chat", "context_length": 4096}]}, + provider="cohere", + ) + malformed_mistral = pcs.resolve_provider( + {"data": [{"id": "future", "capabilities": ["completion_chat"]}]}, + provider="mistral", + ) + + assert malformed_cohere.provider_id == "cohere" assert malformed_cohere.shape_id == "fallback.models.envelope.v1" assert malformed_cohere.fallback is True - assert malformed_mistral.provider_id == pcs.PROVIDER_UNKNOWN + assert malformed_mistral.provider_id == "mistral" assert malformed_mistral.shape_id == "fallback.models.data.v1" assert malformed_mistral.fallback is True +def test_explicit_fallback_and_mixed_native_payloads_are_normalized_per_item(): + malformed = {"models": [{"name": "unsafe", "endpoints": "chat", "context_length": 4096}]} + malformed_record = records_from_payload(malformed, vendor="cohere")[0] + + assert malformed_record.vendor == "cohere" + assert malformed_record.capability.family == mc.FAMILY_UNKNOWN + assert dict(malformed_record.capability.limits) == {} + assert malformed_record.catalog_shape_id == "fallback.models.envelope.v1" + assert malformed_record.fallback is True + + valid_item = { + "name": "native", + "endpoints": ["chat"], + "context_length": 131072, + } + native, fallback = records_from_payload( + {"models": [valid_item, malformed["models"][0]]}, + vendor="cohere", + ) + + assert native.model_id == "native" + assert native.capability.family == mc.FAMILY_CHAT + assert dict(native.capability.limits) == {"context_tokens": 131072} + assert native.catalog_shape_id == "cohere.models.rich.v1" + assert native.fallback is False + assert fallback.model_id == "unsafe" + assert fallback.capability.family == mc.FAMILY_UNKNOWN + assert dict(fallback.capability.limits) == {} + assert fallback.catalog_shape_id == "fallback.models.envelope.v1" + assert fallback.fallback is True + + +def test_provider_specific_reader_is_not_used_for_a_different_fallback_envelope(): + record = records_from_payload( + [{"id": "untrusted", "architecture": {"modality": "text+image->text"}}], + vendor="openrouter", + )[0] + + assert record.vendor == "openrouter" + assert record.capability.family == mc.FAMILY_UNKNOWN + assert record.capability.capabilities == () + assert record.catalog_shape_id == "fallback.models.list.v1" + assert record.fallback is True + + def test_fallback_reader_is_identity_only_even_for_dangerous_looking_fields(): payload = [ { @@ -375,6 +522,41 @@ def test_sglang_model_info_maps_native_generation_flags_only(): assert pooling.capability.family == mc.FAMILY_UNKNOWN +def test_sglang_openai_catalog_preserves_only_valid_native_context_limit(): + valid = records_from_payload( + { + "data": [ + { + "id": "served-model", + "owned_by": "sglang", + "root": "org/model", + "max_model_len": 131072, + } + ] + } + )[0] + nonpositive = sglang.records_from_payload( + { + "data": [ + { + "id": "served-model", + "owned_by": "sglang", + "root": "org/model", + "max_model_len": 0, + } + ] + } + )[0] + + assert valid.vendor == "sglang" + assert valid.capability.family == mc.FAMILY_UNKNOWN + assert valid.capability.capabilities == () + assert dict(valid.capability.limits) == {"context_tokens": 131072} + assert valid.catalog_shape_id == "sglang.models.openai.v1" + assert valid.fallback is False + assert dict(nonpositive.capability.limits) == {} + + def test_identity_only_native_catalogs_remain_unknown(): anthropic_record = anthropic.records_from_payload( { @@ -461,7 +643,8 @@ def test_reader_wrapper_adds_one_lean_evidence_object(): "capabilities": {"completion_chat": True, "function_calling": True}, } ] - } + }, + base_url="https://api.mistral.ai/v1", )[0] serialized = record.to_dict() @@ -472,7 +655,7 @@ def test_reader_wrapper_adds_one_lean_evidence_object(): assert serialized["evidence"] == { "source": mc.SOURCE_PROVIDER_READER, "confidence": mc.CONFIDENCE_PROVIDER_REPORTED, - "provider_source": pcs.PROVIDER_SOURCE_PAYLOAD, + "provider_source": pcs.PROVIDER_SOURCE_HOST, "shape": "mistral.models.rich.v1", "fallback": False, }