main
py 263 lines 10.7 KB
Raw
1 import json
2
3 from helpers.api import ApiHandler, Request, Response
4 from helpers import defer, files, plugins
5 from helpers.extension import call_extensions_async
6 from helpers.persist_chat import _write_atomic, save_tmp_chat
7 from agent import AgentContext
8 from plugins._model_config.api.model_override import _notify_model_override_changed
9 from plugins._model_config.helpers import model_config
10
11
12 def _rename_preset_references(renames: object) -> None:
13 if not isinstance(renames, list):
14 return
15 mapping = {
16 str(item.get("from") or "").strip().casefold(): str(item.get("to") or "").strip()
17 for item in renames
18 if isinstance(item, dict)
19 and str(item.get("from") or "").strip()
20 and str(item.get("to") or "").strip()
21 }
22 mapping.pop(model_config.DEFAULT_PRESET_NAME.casefold(), None)
23 if not mapping:
24 return
25
26 assets = plugins.find_plugin_assets(
27 plugins.CONFIG_FILE_NAME,
28 plugin_name="_model_config",
29 project_name="*",
30 agent_profile="*",
31 only_first=False,
32 )
33 for asset in assets:
34 path = str(asset.get("path") or "")
35 try:
36 config = json.loads(files.read_file(path))
37 except Exception:
38 continue
39 if not isinstance(config, dict):
40 continue
41 selected = str(config.get(model_config.MODEL_PRESET_CONFIG_KEY) or "").strip()
42 replacement = mapping.get(selected.casefold())
43 if replacement:
44 config[model_config.MODEL_PRESET_CONFIG_KEY] = replacement
45 files.write_file(path, json.dumps(config))
46
47 # Chats are lazy-loaded, so update durable references as well as live
48 # AgentContext objects. Otherwise renaming a preset would silently send an
49 # unopened chat back to Default the next time it is loaded.
50 loaded_context_ids = {str(context.id) for context in AgentContext.all()}
51 chat_paths = files.find_existing_paths_by_pattern(
52 files.get_abs_path("usr", "chats", "*", "chat.json")
53 )
54 for path in chat_paths:
55 chat_id = files.basename(files.dirname(path))
56 if chat_id in loaded_context_ids:
57 continue
58 try:
59 chat = json.loads(files.read_file(path))
60 except Exception:
61 continue
62 data = chat.get("data") if isinstance(chat, dict) else None
63 override = data.get("chat_model_override") if isinstance(data, dict) else None
64 selected = str(override.get("preset_name") or "").strip() if isinstance(override, dict) else ""
65 replacement = mapping.get(selected.casefold())
66 if not replacement:
67 continue
68 data["chat_model_override"] = {"preset_name": replacement}
69 _write_atomic(path, json.dumps(chat, ensure_ascii=False))
70
71 for context in AgentContext.all():
72 override = context.get_data("chat_model_override")
73 if not isinstance(override, dict):
74 continue
75 selected = str(override.get("preset_name") or "").strip()
76 replacement = mapping.get(selected.casefold())
77 if not replacement:
78 continue
79 context.set_data("chat_model_override", {"preset_name": replacement})
80 save_tmp_chat(context)
81 _notify_model_override_changed(context)
82
83
84 def _retired_preset_references(
85 previous_names: set[str],
86 saved_names: set[str],
87 ) -> list[dict[str, str]]:
88 saved_by_casefold = {name.casefold(): name for name in saved_names}
89 return [
90 {
91 "from": name,
92 "to": saved_by_casefold.get(
93 name.casefold(),
94 model_config.DEFAULT_PRESET_NAME,
95 ),
96 }
97 for name in previous_names - saved_names
98 ]
99
100
101 def _embedding_signature(preset: dict | None):
102 if not isinstance(preset, dict):
103 return {}
104 default = model_config.resolve_preset(model_config.DEFAULT_PRESET_NAME) or {}
105 config = model_config.preset_to_config(default)
106 if str(preset.get("name") or "") != model_config.DEFAULT_PRESET_NAME:
107 config = model_config.build_config_from_preset(
108 preset,
109 config,
110 strip_api_key=False,
111 )
112 embedding = config.get("embedding_model")
113 return embedding if isinstance(embedding, dict) else {}
114
115
116 def _notify_embedding_changed() -> None:
117 defer.DeferredTask().start_task(call_extensions_async, "embedding_model_changed")
118
119
120 class ModelPresets(ApiHandler):
121 async def process(self, input: dict, request: Request) -> dict | Response:
122 action = input.get("action", "get")
123 project_name = str(input.get("project_name") or "").strip()
124 agent_profile = str(input.get("agent_profile") or "").strip()
125 context_id = str(input.get("context_id") or "").strip()
126 scope = str(input.get("scope") or "").strip()
127
128 if action == "get":
129 if scope == "project":
130 return {"ok": True, "presets": []}
131 presets = model_config.get_presets()
132 context = AgentContext.get(context_id) if context_id else None
133 if context:
134 configured = model_config.get_configured_preset_name(agent=context.agent0)
135 selected = model_config.get_effective_preset_name(context.agent0)
136 else:
137 configured = model_config.get_configured_preset_name(
138 project_name=project_name or None,
139 agent_profile=agent_profile or None,
140 )
141 selected = configured
142 return {
143 "ok": True,
144 "presets": presets,
145 "global_presets": presets,
146 "project_presets": [],
147 "configured_preset": configured,
148 "selected_preset": selected,
149 }
150
151 elif action == "save":
152 presets = input.get("presets")
153 if not isinstance(presets, list):
154 return Response(status=400, response="presets must be an array")
155 if scope == "project" or project_name:
156 return Response(
157 status=400,
158 response="Preset definitions are global; select a global preset for this project.",
159 )
160 previous_names = {
161 str(preset.get("name") or "") for preset in model_config.get_presets()
162 }
163 previous_embeddings = {
164 str(preset.get("name") or ""): _embedding_signature(preset)
165 for preset in model_config.get_presets()
166 }
167 try:
168 model_config.save_presets(presets)
169 except ValueError as exc:
170 return Response(status=400, response=str(exc))
171 saved_names = {
172 str(preset.get("name") or "") for preset in model_config.get_presets()
173 }
174 retired = _retired_preset_references(previous_names, saved_names)
175 renames = input.get("renames") if isinstance(input.get("renames"), list) else []
176 _rename_preset_references([*retired, *renames])
177 saved_embeddings = {
178 str(preset.get("name") or ""): _embedding_signature(preset)
179 for preset in model_config.get_presets()
180 }
181 if previous_embeddings != saved_embeddings:
182 _notify_embedding_changed()
183 return {"ok": True, "presets": model_config.get_presets()}
184
185 elif action == "reset":
186 if scope == "project" or project_name:
187 return Response(status=400, response="Project presets cannot be reset.")
188 previous = model_config.get_presets()
189 previous_embeddings = {
190 str(preset.get("name") or ""): _embedding_signature(preset)
191 for preset in previous
192 }
193 presets = model_config.reset_presets()
194 saved_names = {str(preset.get("name") or "") for preset in presets}
195 previous_names = {str(preset.get("name") or "") for preset in previous}
196 retired = _retired_preset_references(previous_names, saved_names)
197 _rename_preset_references(retired)
198 current_embeddings = {
199 str(preset.get("name") or ""): _embedding_signature(preset)
200 for preset in presets
201 }
202 if previous_embeddings != current_embeddings:
203 _notify_embedding_changed()
204 return {"ok": True, "presets": presets}
205
206 elif action == "select":
207 name = str(input.get("name") or "").strip()
208 preset = model_config.resolve_preset(name)
209 if not preset:
210 return Response(status=404, response=f"Preset '{name}' not found")
211 canonical_name = str(preset.get("name") or model_config.DEFAULT_PRESET_NAME)
212 context = AgentContext.get(context_id) if context_id else None
213 previous_embedding = (
214 model_config.get_embedding_model_config(context.agent0)
215 if context
216 else None
217 )
218 previous_name = (
219 model_config.get_effective_preset_name(context.agent0)
220 if context
221 else model_config.get_configured_preset_name(
222 project_name=project_name or None,
223 agent_profile=agent_profile or None,
224 )
225 )
226 previous = model_config.resolve_preset(previous_name)
227 plugins.save_plugin_config(
228 "_model_config",
229 project_name,
230 agent_profile,
231 {model_config.MODEL_PRESET_CONFIG_KEY: canonical_name},
232 )
233
234 if context:
235 context.set_data("chat_model_override", {"preset_name": canonical_name})
236 save_tmp_chat(context)
237 _notify_model_override_changed(context)
238 current_embedding = (
239 model_config.get_embedding_model_config(context.agent0)
240 if context
241 else _embedding_signature(preset)
242 )
243 if (previous_embedding or _embedding_signature(previous)) != current_embedding:
244 _notify_embedding_changed()
245 return {"ok": True, "selected_preset": canonical_name}
246
247 elif action == "resolve":
248 name = str(input.get("name") or "").strip()
249 if not name:
250 return Response(status=400, response="Missing preset name")
251 resolved = model_config.resolve_preset(name)
252 if not resolved:
253 return Response(status=404, response=f"Preset '{name}' not found")
254 return {
255 "ok": True,
256 "preset": {
257 **resolved,
258 "scope": "global",
259 "project_name": "",
260 },
261 }
262
263 return Response(status=400, response=f"Unknown action: {action}")