main
py 105 lines 4.42 KB
Raw
1 import time
2
3 from helpers.api import ApiHandler, Request, Response
4 from helpers import defer
5 from helpers.extension import call_extensions_async
6 from helpers.persist_chat import save_tmp_chat
7 from agent import AgentContext
8 from plugins._model_config.helpers import model_config
9
10 _MODEL_OVERRIDE_REVISION_KEY = "_model_config_override_revision"
11
12
13 def _notify_model_override_changed(ctx: AgentContext) -> None:
14 ctx.set_output_data(_MODEL_OVERRIDE_REVISION_KEY, time.time())
15
16 try:
17 from helpers.state_monitor_integration import mark_dirty_for_context
18
19 mark_dirty_for_context(ctx.id, reason="model_config.model_override")
20 except Exception:
21 pass
22
23
24 def _notify_embedding_if_changed(before: dict, after: dict) -> None:
25 if before != after:
26 defer.DeferredTask().start_task(call_extensions_async, "embedding_model_changed")
27
28
29 class ModelOverride(ApiHandler):
30 async def process(self, input: dict, request: Request) -> dict | Response:
31 context_id = input.get("context_id", "")
32 action = input.get("action", "get") # get | set | set_preset | clear
33
34 if not context_id:
35 return Response(status=400, response="Missing context_id")
36
37 ctx = AgentContext.get(context_id)
38 if not ctx:
39 return Response(status=404, response="Context not found")
40
41 if action == "get":
42 override = ctx.get_data("chat_model_override")
43 allowed = model_config.is_chat_override_allowed(ctx.agent0)
44 return {
45 "override": override,
46 "allowed": allowed,
47 "configured_preset": model_config.get_configured_preset_name(agent=ctx.agent0),
48 "effective_preset": model_config.get_effective_preset_name(ctx.agent0),
49 }
50
51 elif action == "set":
52 if not model_config.is_chat_override_allowed(ctx.agent0):
53 return Response(status=403, response="Per-chat override is disabled")
54 override_config = input.get("override")
55 if not override_config or not isinstance(override_config, dict):
56 return Response(status=400, response="Missing or invalid override config")
57 previous_embedding = model_config.get_embedding_model_config(ctx.agent0)
58 ctx.set_data("chat_model_override", override_config)
59 save_tmp_chat(ctx)
60 _notify_model_override_changed(ctx)
61 _notify_embedding_if_changed(
62 previous_embedding,
63 model_config.get_embedding_model_config(ctx.agent0),
64 )
65 return {"ok": True, "override": override_config}
66
67 elif action == "set_preset":
68 if not model_config.is_chat_override_allowed(ctx.agent0):
69 return Response(status=403, response="Per-chat override is disabled")
70 preset_name = input.get("preset_name", "")
71 if not preset_name:
72 return Response(status=400, response="Missing preset_name")
73 previous_embedding = model_config.get_embedding_model_config(ctx.agent0)
74 # Verify preset exists
75 preset = model_config.get_preset_by_name(preset_name)
76 if not preset:
77 return Response(status=404, response=f"Preset '{preset_name}' not found")
78 # Store as a preset reference
79 canonical_name = str(preset.get("name") or preset_name)
80 override_value = {"preset_name": canonical_name}
81 ctx.set_data("chat_model_override", override_value)
82 save_tmp_chat(ctx)
83 _notify_model_override_changed(ctx)
84 _notify_embedding_if_changed(
85 previous_embedding,
86 model_config.get_embedding_model_config(ctx.agent0),
87 )
88 return {"ok": True, "preset_name": canonical_name}
89
90 elif action == "clear":
91 previous_embedding = model_config.get_embedding_model_config(ctx.agent0)
92 ctx.set_data("chat_model_override", None)
93 save_tmp_chat(ctx)
94 _notify_model_override_changed(ctx)
95 _notify_embedding_if_changed(
96 previous_embedding,
97 model_config.get_embedding_model_config(ctx.agent0),
98 )
99 return {
100 "ok": True,
101 "override": None,
102 "effective_preset": model_config.get_configured_preset_name(agent=ctx.agent0),
103 }
104
105 return Response(status=400, response=f"Unknown action: {action}")