fix(models): harden capability catalog normalization

This commit is contained in:
RaresKeY 2026-07-18 11:20:19 +00:00
parent 47e68b4eb1
commit d1ad5de108
8 changed files with 496 additions and 71 deletions

27
app.py
View file

@ -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)

View file

@ -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.

View file

@ -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,

View file

@ -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)

View file

@ -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",

View file

@ -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))

View file

@ -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

View file

@ -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,
}