-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmodel_config.py
More file actions
132 lines (106 loc) · 3.75 KB
/
Copy pathmodel_config.py
File metadata and controls
132 lines (106 loc) · 3.75 KB
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
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
"""Observer / Reflector model configuration — user-facing API.
Tiny module: just dispatch into ``model_presets`` (data) and
``server_manager`` (storage). Validation rules live here.
"""
from __future__ import annotations
import json
import re
from typing import Any
try:
from .model_presets import ( # type: ignore
DEFAULT_OBSERVER,
DEFAULT_REFLECTOR,
PRESETS,
find_preset,
)
from .server_manager import config_path, save_config
except ImportError:
from model_presets import ( # type: ignore[no-redef]
DEFAULT_OBSERVER,
DEFAULT_REFLECTOR,
PRESETS,
find_preset,
)
from server_manager import config_path, save_config # type: ignore[no-redef]
ROLES = ("observer", "reflector")
_ENV_VAR_RE = re.compile(r"^[A-Z][A-Z0-9_]*$")
def _load_user_raw() -> dict:
p = config_path()
if not p.exists():
return {}
try:
data = json.loads(p.read_text(encoding="utf-8"))
return data if isinstance(data, dict) else {}
except Exception:
return {}
def _check_role(role: str) -> None:
if role not in ROLES:
raise ValueError(f"unknown role '{role}' (valid: {', '.join(ROLES)})")
def _check_name(name: str) -> None:
if not name or not name.strip():
raise ValueError("model name must be a non-empty string")
def _check_base_url(url: str) -> None:
if not url.startswith(("http://", "https://")):
raise ValueError(f"base_url must start with http:// or https:// (got '{url}')")
def _check_api_key_env(env_var: str) -> None:
if env_var == "":
return
if not _ENV_VAR_RE.match(env_var):
raise ValueError(
f"api_key_env must be UPPER_SNAKE_CASE matching ^[A-Z][A-Z0-9_]*$ (got '{env_var}')"
)
def _resolve_field(role: str, field: str, role_user: dict, legacy: dict, default: dict) -> str:
for src in (role_user, legacy, default):
v = src.get(field)
if v:
return v
return ""
def get_model(role: str) -> dict[str, str]:
"""Return ``{name, base_url, api_key_env}`` for the requested role.
Resolution: explicit ``<role>_*`` keys → (Observer only) legacy ``model_*``
→ built-in defaults.
"""
_check_role(role)
user_raw = _load_user_raw()
role_user = {
"name": user_raw.get(f"{role}_name"),
"base_url": user_raw.get(f"{role}_url"),
"api_key_env": user_raw.get(f"{role}_api_key_env"),
}
legacy = (
{
"name": user_raw.get("model_name"),
"base_url": user_raw.get("model_url"),
"api_key_env": user_raw.get("model_api_key_env"),
}
if role == "observer"
else {}
)
default = DEFAULT_OBSERVER if role == "observer" else DEFAULT_REFLECTOR
return {
"name": _resolve_field(role, "name", role_user, legacy, default),
"base_url": _resolve_field(role, "base_url", role_user, legacy, default),
"api_key_env": _resolve_field(role, "api_key_env", role_user, legacy, default),
}
def set_model(role: str, *, name: str, base_url: str, api_key_env: str) -> None:
_check_role(role)
_check_name(name)
_check_base_url(base_url)
_check_api_key_env(api_key_env)
save_config(
{
f"{role}_name": name,
f"{role}_url": base_url,
f"{role}_api_key_env": api_key_env,
}
)
def list_presets() -> list[dict[str, Any]]:
return [dict(p) for p in PRESETS]
def apply_preset(preset_id: str) -> dict[str, Any]:
p = find_preset(preset_id)
if not p:
valid = ", ".join(x["id"] for x in PRESETS)
raise ValueError(f"unknown preset '{preset_id}' (valid: {valid})")
set_model("observer", **p["observer"])
set_model("reflector", **p["reflector"])
return p