145 lines
5.3 KiB
Python
145 lines
5.3 KiB
Python
from __future__ import annotations
|
|
|
|
import os
|
|
from dataclasses import dataclass
|
|
from pathlib import Path
|
|
|
|
from dotenv import load_dotenv
|
|
|
|
|
|
BACKEND_ROOT = Path(__file__).resolve().parents[1]
|
|
PROJECT_ROOT = BACKEND_ROOT.parent
|
|
|
|
load_dotenv(BACKEND_ROOT / ".env")
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ProviderModel:
|
|
id: str
|
|
vision: bool = False
|
|
# Strict function schemas are provider/model capabilities, not an
|
|
# assumption about every OpenAI-compatible endpoint.
|
|
strict_tool_schema: bool = False
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ProviderConfig:
|
|
id: str
|
|
label: str
|
|
base_url: str
|
|
api_key: str
|
|
models: tuple[ProviderModel, ...]
|
|
|
|
@property
|
|
def configured(self) -> bool:
|
|
return bool(self.base_url and self.api_key and self.models)
|
|
|
|
def model(self, model_id: str) -> ProviderModel | None:
|
|
return next((model for model in self.models if model.id == model_id), None)
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class Settings:
|
|
task_root: Path
|
|
conversation_root: Path
|
|
library_root: Path
|
|
engine_root: Path
|
|
llm_base_url: str
|
|
llm_api_key: str
|
|
llm_model: str
|
|
llm_timeout_s: float
|
|
default_provider_id: str
|
|
providers: tuple[ProviderConfig, ...]
|
|
max_repair_attempts: int = 4
|
|
|
|
@property
|
|
def llm_configured(self) -> bool:
|
|
return self.provider_for(self.default_provider_id) is not None
|
|
|
|
def provider_for(self, provider_id: str | None) -> ProviderConfig | None:
|
|
requested = str(provider_id or self.default_provider_id).strip().lower()
|
|
return next((provider for provider in self.providers if provider.id == requested and provider.configured), None)
|
|
|
|
def resolve_model(self, provider_id: str | None, model_id: str | None) -> tuple[ProviderConfig, ProviderModel]:
|
|
provider = self.provider_for(provider_id)
|
|
if provider is None:
|
|
raise ValueError("The selected model provider is not configured")
|
|
selected = str(model_id or "").strip() or provider.models[0].id
|
|
model = provider.model(selected)
|
|
if model is None:
|
|
raise ValueError("The selected model is not enabled for this provider")
|
|
return provider, model
|
|
|
|
|
|
def _enabled_model_ids(value: str) -> set[str]:
|
|
return {item.strip() for item in value.split(",") if item.strip()}
|
|
|
|
|
|
def _as_bool(value: str) -> bool:
|
|
return value.strip().lower() in {"1", "true", "yes", "on"}
|
|
|
|
|
|
def _models(
|
|
value: str,
|
|
vision_value: str = "",
|
|
strict_value: str = "",
|
|
strict_all: bool = False,
|
|
) -> tuple[ProviderModel, ...]:
|
|
vision_ids = {item.strip() for item in vision_value.split(",") if item.strip()}
|
|
strict_ids = _enabled_model_ids(strict_value)
|
|
return tuple(
|
|
ProviderModel(
|
|
id=item,
|
|
vision=item in vision_ids,
|
|
strict_tool_schema=strict_all or item in strict_ids,
|
|
)
|
|
for item in (part.strip() for part in value.split(","))
|
|
if item
|
|
)
|
|
|
|
|
|
def _provider(prefix: str, provider_id: str, label: str, default_base_url: str, default_model: str = "") -> ProviderConfig:
|
|
# The legacy CDSL_LLM_* variables remain the DeepSeek default so existing
|
|
# local installations continue to work without copying secrets.
|
|
legacy = provider_id == "deepseek"
|
|
base_url = os.getenv(f"CDSL_{prefix}_BASE_URL", os.getenv("CDSL_LLM_BASE_URL", default_base_url) if legacy else default_base_url).rstrip("/")
|
|
api_key = os.getenv(f"CDSL_{prefix}_API_KEY", os.getenv("CDSL_LLM_API_KEY", "") if legacy else "")
|
|
model_list = os.getenv(f"CDSL_{prefix}_MODELS", os.getenv("CDSL_LLM_MODEL", default_model) if legacy else default_model)
|
|
vision_models = os.getenv(f"CDSL_{prefix}_VISION_MODELS", "")
|
|
# A provider-wide switch is convenient for a verified endpoint. The model
|
|
# list lets mixed capability deployments opt in only selected models.
|
|
strict_all = _as_bool(os.getenv(f"CDSL_{prefix}_STRICT_TOOL_SCHEMA", ""))
|
|
strict_models = os.getenv(f"CDSL_{prefix}_STRICT_TOOL_MODELS", "")
|
|
return ProviderConfig(
|
|
provider_id,
|
|
label,
|
|
base_url,
|
|
api_key,
|
|
_models(model_list, vision_models, strict_models, strict_all),
|
|
)
|
|
|
|
|
|
def get_settings() -> Settings:
|
|
data_root = BACKEND_ROOT / "data"
|
|
providers = (
|
|
_provider("DEEPSEEK", "deepseek", "DeepSeek", "https://api.deepseek.com/v1", "deepseek-chat"),
|
|
_provider("OPENAI", "openai", "OpenAI", "https://api.openai.com/v1"),
|
|
_provider("KIMI", "kimi", "Kimi", "https://api.moonshot.cn/v1"),
|
|
)
|
|
default_provider_id = os.getenv("CDSL_DEFAULT_PROVIDER", "deepseek").strip().lower() or "deepseek"
|
|
default_provider = next((item for item in providers if item.id == default_provider_id), providers[0])
|
|
default_model = os.getenv("CDSL_DEFAULT_MODEL", "").strip() or (default_provider.models[0].id if default_provider.models else "")
|
|
return Settings(
|
|
task_root=data_root / "tasks",
|
|
conversation_root=data_root / "conversations",
|
|
library_root=BACKEND_ROOT / "cdsl_library",
|
|
engine_root=BACKEND_ROOT / "engine" / "cdsl_engine",
|
|
llm_base_url=default_provider.base_url,
|
|
llm_api_key=default_provider.api_key,
|
|
llm_model=default_model,
|
|
llm_timeout_s=float(os.getenv("CDSL_LLM_TIMEOUT_S", "90")),
|
|
max_repair_attempts=max(0, int(os.getenv("CDSL_MAX_REPAIR_ATTEMPTS", "4"))),
|
|
default_provider_id=default_provider_id,
|
|
providers=providers,
|
|
)
|