Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 10 additions & 1 deletion agent_fastapi.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -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:
Expand Down
44 changes: 44 additions & 0 deletions src/open_storyline/model_presets.py
Original file line number Diff line number Diff line change
@@ -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
51 changes: 51 additions & 0 deletions tests/test_model_presets.py
Original file line number Diff line number Diff line change
@@ -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()