fix: preserve traceback and avoid mutating llm clients

This commit is contained in:
rookiestar28
2026-03-20 12:16:23 +08:00
parent bf2345310e
commit 7e0c1ee074
9 changed files with 147 additions and 38 deletions
+1 -1
View File
@@ -531,7 +531,7 @@ async def llm_models_handler(request: web.Request) -> web.Response:
return web.json_response(
{"ok": False, "error": f"Upstream error: {str_e}"}, status=502
)
raise e
raise
except Exception as e:
stale = get_stale_cached_models(target.cache_key)
+4 -4
View File
@@ -40,7 +40,7 @@ class OpenClawImageToPrompt:
# Refresh the default LLMClient per call so provider/key changes from Settings/UI
# apply without restarting ComfyUI. Preserve injected mocks/fakes for tests.
if isinstance(self.llm_client, LLMClient):
self.llm_client = LLMClient()
return LLMClient()
return self.llm_client
@classmethod
@@ -132,10 +132,10 @@ Do not use markdown blocks.
return (caption, tags_str, prompt_suggestion)
except Exception as e:
except Exception:
metrics.increment("errors")
logger.error(f"Failed to generate prompt from image: {e}")
raise e
logger.error("Failed to generate prompt from image", exc_info=True)
raise
# IMPORTANT: keep legacy class alias for existing imports and tests.
+1 -1
View File
@@ -47,7 +47,7 @@ class AutomationComposerService:
# Refresh the default LLMClient per request so UI-saved provider/key updates
# take effect without backend restart. Keep injected fakes intact for tests.
if isinstance(self.llm_client, LLMClient):
self.llm_client = LLMClient()
return LLMClient()
return self.llm_client
def consume_last_reasoning_debug(self) -> Any:
+6 -4
View File
@@ -45,8 +45,10 @@ class PlannerService:
# Assist handlers are long-lived singletons, so caching the startup client here
# causes stale provider/key state after UI Save (requires backend restart).
# Keep custom test fakes/injected clients intact by only rotating real LLMClient.
# IMPORTANT: do not overwrite self.llm_client here; mutating shared service state
# just to obtain a fresh default client is needlessly racy under concurrent access.
if isinstance(self.llm_client, LLMClient):
self.llm_client = LLMClient()
return LLMClient()
return self.llm_client
def consume_last_reasoning_debug(self) -> Any:
@@ -197,7 +199,7 @@ Style: {style_directives}
return positive, negative, validated_params.dict()
except Exception as e:
except Exception:
metrics.increment("errors")
logger.error(f"Failed to plan generation: {e}")
raise e
logger.error("Failed to plan generation", exc_info=True)
raise
+6 -4
View File
@@ -51,8 +51,10 @@ class RefinerService:
# Refiner shares the same long-lived assist handler lifecycle as Planner; keeping
# the startup client causes stale provider/key state after UI Save.
# Preserve injected fakes by only rotating real LLMClient instances.
# IMPORTANT: do not mutate the stored long-lived service client when resolving
# a fresh default request client; that write is unnecessary shared-state churn.
if isinstance(self.llm_client, LLMClient):
self.llm_client = LLMClient()
return LLMClient()
return self.llm_client
def consume_last_reasoning_debug(self) -> Any:
@@ -254,7 +256,7 @@ Issue: {issue}
return refined_pos, refined_neg, final_patch, rationale
except Exception as e:
except Exception:
metrics.increment("errors")
logger.error(f"Refiner failed: {e}")
raise e
logger.error("Refiner failed", exc_info=True)
raise
+3 -2
View File
@@ -10,9 +10,10 @@ function resolveUiTimeoutMs() {
}
// IMPORTANT: WSL on /mnt/* can load the module-heavy harness much slower than
// native filesystems; give readiness checks extra budget to avoid false reds.
// native filesystems, especially after several sequential page reloads in the
// same worker; give readiness checks extra budget to avoid false reds.
if (process.platform === 'linux' && process.env.WSL_DISTRO_NAME && process.cwd().startsWith('/mnt/')) {
return 60_000;
return 120_000;
}
return 30_000;
+6 -10
View File
@@ -80,16 +80,14 @@ class TestAssistLLMClientHotReload(unittest.TestCase):
_PlannerDynamicFakeLLMClient._next_id = 0
with patch.object(planner_mod, "LLMClient", _PlannerDynamicFakeLLMClient):
planner = planner_mod.PlannerService()
init_client = planner.llm_client
pos1, _, _ = planner.plan_generation("SDXL-v1", "req", "style", seed=1)
first_client = planner.llm_client
pos2, _, _ = planner.plan_generation("SDXL-v1", "req", "style", seed=2)
second_client = planner.llm_client
self.assertNotEqual(pos1, pos2)
self.assertIsNot(first_client, second_client)
self.assertEqual(first_client.instance_id, 2)
self.assertEqual(second_client.instance_id, 3)
self.assertIs(planner.llm_client, init_client)
self.assertEqual(init_client.instance_id, 1)
def test_refiner_refreshes_default_llm_client_per_request(self):
import services.refiner as refiner_mod
@@ -97,6 +95,7 @@ class TestAssistLLMClientHotReload(unittest.TestCase):
_RefinerDynamicFakeLLMClient._next_id = 0
with patch.object(refiner_mod, "LLMClient", _RefinerDynamicFakeLLMClient):
refiner = refiner_mod.RefinerService()
init_client = refiner.llm_client
rp1, _, _, _ = refiner.refine_prompt(
image_b64="dummy",
@@ -105,7 +104,6 @@ class TestAssistLLMClientHotReload(unittest.TestCase):
issue="fix",
params_json=json.dumps({"width": 1024, "height": 1024}),
)
first_client = refiner.llm_client
rp2, _, _, _ = refiner.refine_prompt(
image_b64="dummy",
orig_positive="op",
@@ -113,12 +111,10 @@ class TestAssistLLMClientHotReload(unittest.TestCase):
issue="fix",
params_json=json.dumps({"width": 1024, "height": 1024}),
)
second_client = refiner.llm_client
self.assertNotEqual(rp1, rp2)
self.assertIsNot(first_client, second_client)
self.assertEqual(first_client.instance_id, 2)
self.assertEqual(second_client.instance_id, 3)
self.assertIs(refiner.llm_client, init_client)
self.assertEqual(init_client.instance_id, 1)
def test_planner_keeps_injected_custom_llm_client(self):
from services.planner import PlannerService
+5 -12
View File
@@ -51,7 +51,7 @@ class TestLLMClientHotReloadNonAssist(unittest.TestCase):
patch.dict(os.environ, {"OPENCLAW_ENABLE_TOOL_CALLING": "1"}),
):
svc = composer_mod.AutomationComposerService()
first_init_client = svc.llm_client
init_client = svc.llm_client
res1 = svc.compose_payload(
kind="trigger",
@@ -59,19 +59,15 @@ class TestLLMClientHotReloadNonAssist(unittest.TestCase):
intent="compose 1",
inputs_hint={"requirements": "a"},
)
first_request_client = svc.llm_client
res2 = svc.compose_payload(
kind="trigger",
template_id="tmpl",
intent="compose 2",
inputs_hint={"requirements": "b"},
)
second_request_client = svc.llm_client
self.assertIsNot(first_init_client, first_request_client)
self.assertIsNot(first_request_client, second_request_client)
self.assertEqual(first_request_client.instance_id, 2)
self.assertEqual(second_request_client.instance_id, 3)
self.assertIs(svc.llm_client, init_client)
self.assertEqual(init_client.instance_id, 1)
self.assertFalse(res1["used_tool_calling"])
self.assertFalse(res2["used_tool_calling"])
self.assertTrue(any("tool_call_fallback" in w for w in res1["warnings"]))
@@ -98,12 +94,9 @@ class TestLLMClientHotReloadNonAssist(unittest.TestCase):
detail_level="medium",
max_image_side=512,
)
second_request_client = node.llm_client
self.assertIsNot(init_client, first_request_client)
self.assertIsNot(first_request_client, second_request_client)
self.assertEqual(first_request_client.instance_id, 2)
self.assertEqual(second_request_client.instance_id, 3)
self.assertIs(node.llm_client, init_client)
self.assertEqual(init_client.instance_id, 1)
self.assertNotEqual(cap1, cap2)
self.assertEqual(tags1, "tag1, tag2")
self.assertNotEqual(prompt1, prompt2)
+115
View File
@@ -0,0 +1,115 @@
import traceback
import unittest
from unittest.mock import MagicMock, patch
class _BoomLLMClient:
def complete(self, *args, **kwargs):
raise RuntimeError("boom")
class TestR155ExceptionFidelity(unittest.TestCase):
def _assert_traceback_contains_frame(
self, tb, filename_fragment, expected_lineno, expected_line_fragment
):
frames = traceback.extract_tb(tb)
self.assertTrue(
any(
filename_fragment in frame.filename.replace("\\", "/")
and frame.lineno == expected_lineno
and expected_line_fragment in (frame.line or "")
for frame in frames
),
msg="Traceback frames did not include the original failing call site:\n"
+ "\n".join(
f"{frame.filename}:{frame.lineno}: {frame.line}" for frame in frames
),
)
def _capture_runtime_error(self, func):
try:
func()
except RuntimeError as exc:
return exc, exc.__traceback__
self.fail("RuntimeError was not raised")
def test_planner_preserves_original_traceback_line(self):
from services.planner import PlannerService
planner = PlannerService()
planner.llm_client = _BoomLLMClient()
_, tb = self._capture_runtime_error(
lambda: planner.plan_generation("SDXL-v1", "req", "style", seed=1)
)
self._assert_traceback_contains_frame(
tb, "services/planner.py", 160, "response = llm_client.complete("
)
def test_refiner_preserves_original_traceback_line(self):
from services.refiner import RefinerService
refiner = RefinerService()
refiner.llm_client = _BoomLLMClient()
_, tb = self._capture_runtime_error(
lambda: refiner.refine_prompt(
image_b64="dummy",
orig_positive="op",
orig_negative="on",
issue="fix",
params_json="{}",
)
)
self._assert_traceback_contains_frame(
tb, "services/refiner.py", 203, "response = llm_client.complete("
)
def test_image_to_prompt_preserves_original_traceback_line(self):
from nodes.image_to_prompt import OpenClawImageToPrompt
node = OpenClawImageToPrompt()
node.llm_client = _BoomLLMClient()
with patch.object(node, "_tensor_to_base64_png", return_value="ZmFrZQ=="):
_, tb = self._capture_runtime_error(
lambda: node.generate_prompt(
image=object(),
goal="goal",
detail_level="medium",
max_image_side=512,
)
)
self._assert_traceback_contains_frame(
tb, "nodes/image_to_prompt.py", 112, "response = llm_client.complete("
)
def test_api_config_runtime_error_preserves_original_traceback_line(self):
from api.config import llm_models_handler
request = MagicMock()
request.query = {}
request.remote = "127.0.0.1"
with (
patch("api.config.check_rate_limit", return_value=True),
patch("api.config.require_admin_token", return_value=(True, None)),
patch("api.config.get_effective_config", return_value=({"provider": "openai"}, {})),
patch("services.providers.keys.get_api_key_for_provider", return_value="sk-test"),
patch("api.config.fetch_remote_model_list", side_effect=RuntimeError("boom")),
):
_, tb = self._capture_runtime_error(
lambda: self._run_async(llm_models_handler(request))
)
self._assert_traceback_contains_frame(
tb, "api/config.py", 490, "models = fetch_remote_model_list("
)
def _run_async(self, coro):
import asyncio
return asyncio.run(coro)