| 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}") |