diff --git a/agent_fastapi.py b/agent_fastapi.py index ae9ebab..c179b65 100644 --- a/agent_fastapi.py +++ b/agent_fastapi.py @@ -54,6 +54,7 @@ ) from open_storyline.config import load_settings, default_config_path from open_storyline.config import Settings +from open_storyline.model_presets import list_model_presets, resolve_model_preset from open_storyline.storage.agent_memory import ArtifactStore from open_storyline.mcp.hooks.node_interceptors import ToolInterceptor from open_storyline.mcp.hooks.chat_middleware import set_mcp_log_sink, reset_mcp_log_sink @@ -1274,7 +1275,7 @@ def __init__(self, session_id: str, cfg: Settings): default_llm = _peek_builtin_model_name("llm", self.cfg) default_vlm = _peek_builtin_model_name("vlm", self.cfg) - self.chat_models = [default_llm, CUSTOM_MODEL_KEY] + self.chat_models = [default_llm, *list_model_presets("llm"), CUSTOM_MODEL_KEY] self.chat_model_key = default_llm self.vlm_models = [default_vlm, CUSTOM_MODEL_KEY] @@ -1934,6 +1935,14 @@ async def ensure_agent(self) -> None: if not isinstance(self.custom_llm_config, dict): raise RuntimeError("please fill in model/base_url/api_key of custom LLM") llm_override = self.custom_llm_config + elif self.chat_model_key in list_model_presets("llm"): + llm_override, err = resolve_model_preset("llm", self.chat_model_key) + if err: + raise RuntimeError(err) + for key in ("timeout", "temperature", "max_retries", "top_p", "max_tokens"): + value = getattr(self.cfg.llm, key, None) + if value not in (None, ""): + llm_override[key] = value else: llm_override, err = _resolve_builtin_model_override("llm", self.cfg.llm) if err: diff --git a/src/open_storyline/model_presets.py b/src/open_storyline/model_presets.py new file mode 100644 index 0000000..c6b6139 --- /dev/null +++ b/src/open_storyline/model_presets.py @@ -0,0 +1,44 @@ +from __future__ import annotations + +import os +from typing import Any, Dict, Mapping, Optional, Tuple + + +ATLAS_CLOUD_MODEL_KEY = "Atlas Cloud" + +_MODEL_PRESETS = { + "llm": { + ATLAS_CLOUD_MODEL_KEY: { + "model": "deepseek-ai/deepseek-v4-pro", + "base_url": "https://api.atlascloud.ai/v1", + "api_key_env": ("ATLASCLOUD_API_KEY", "ATLAS_CLOUD_API_KEY"), + }, + }, +} + + +def list_model_presets(kind: str) -> list[str]: + return list(_MODEL_PRESETS.get(kind.strip().lower(), {})) + + +def resolve_model_preset( + kind: str, + key: str, + environ: Optional[Mapping[str, str]] = None, +) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: + kind = kind.strip().lower() + preset = _MODEL_PRESETS.get(kind, {}).get(key) + if preset is None: + return None, f"unknown {kind} model preset: {key}" + + env = os.environ if environ is None else environ + api_key_env = preset["api_key_env"] + api_key = next((env.get(name, "").strip() for name in api_key_env if env.get(name, "").strip()), "") + if not api_key: + return None, f"{key} requires one of: {', '.join(api_key_env)}" + + return { + "model": preset["model"], + "base_url": preset["base_url"], + "api_key": api_key, + }, None diff --git a/tests/test_model_presets.py b/tests/test_model_presets.py new file mode 100644 index 0000000..7141d51 --- /dev/null +++ b/tests/test_model_presets.py @@ -0,0 +1,51 @@ +import unittest + +from open_storyline.model_presets import ( + ATLAS_CLOUD_MODEL_KEY, + list_model_presets, + resolve_model_preset, +) + + +class ModelPresetTests(unittest.TestCase): + def test_atlas_cloud_is_available_for_llm_only(self): + self.assertEqual(list_model_presets("llm"), [ATLAS_CLOUD_MODEL_KEY]) + self.assertEqual(list_model_presets("vlm"), []) + + def test_atlas_cloud_resolves_openai_compatible_config(self): + config, error = resolve_model_preset( + "llm", + ATLAS_CLOUD_MODEL_KEY, + {"ATLASCLOUD_API_KEY": "test-key"}, + ) + + self.assertIsNone(error) + self.assertEqual( + config, + { + "model": "deepseek-ai/deepseek-v4-pro", + "base_url": "https://api.atlascloud.ai/v1", + "api_key": "test-key", + }, + ) + + def test_atlas_cloud_accepts_legacy_env_name(self): + config, error = resolve_model_preset( + "llm", + ATLAS_CLOUD_MODEL_KEY, + {"ATLAS_CLOUD_API_KEY": "legacy-key"}, + ) + + self.assertIsNone(error) + self.assertEqual(config["api_key"], "legacy-key") + + def test_atlas_cloud_reports_missing_api_key(self): + config, error = resolve_model_preset("llm", ATLAS_CLOUD_MODEL_KEY, {}) + + self.assertIsNone(config) + self.assertIn("ATLASCLOUD_API_KEY", error) + self.assertIn("ATLAS_CLOUD_API_KEY", error) + + +if __name__ == "__main__": + unittest.main()