Spaces:
Sleeping
Sleeping
File size: 1,494 Bytes
2415446 a1bab2d 2415446 c817fe8 2415446 a1bab2d 2415446 c817fe8 2415446 c817fe8 2415446 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 | """Provider-prefixed model reference helpers."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Protocol
@dataclass(frozen=True, slots=True)
class ConfiguredChatModelRef:
"""A unique configured chat model reference."""
model_ref: str
provider_id: str
model_id: str
class ChatModelConfig(Protocol):
model: str
model_fable: str | None
model_opus: str | None
model_sonnet: str | None
model_haiku: str | None
def parse_provider_type(model_ref: str) -> str:
"""Extract provider type from any 'provider/model' string."""
return model_ref.split("/", 1)[0]
def parse_model_name(model_ref: str) -> str:
"""Extract model name from any 'provider/model' string."""
return model_ref.split("/", 1)[1]
def configured_chat_model_refs(
settings: ChatModelConfig,
) -> tuple[ConfiguredChatModelRef, ...]:
"""Return unique configured chat provider/model refs."""
model_refs = dict.fromkeys(
model_ref
for model_ref in (
settings.model,
settings.model_fable,
settings.model_opus,
settings.model_sonnet,
settings.model_haiku,
)
if model_ref is not None
)
return tuple(
ConfiguredChatModelRef(
model_ref=model_ref,
provider_id=parse_provider_type(model_ref),
model_id=parse_model_name(model_ref),
)
for model_ref in model_refs
)
|