From 1ca571131d769436f116a5da5c17558d9a5396f5 Mon Sep 17 00:00:00 2001 From: liuwanwan1 Date: Tue, 14 Jul 2026 18:31:25 +0800 Subject: [PATCH] fix: resolve priority reliability issues --- .../agent/runners/tool_loop_agent_runner.py | 42 +++++++++++----- .../method/agent_sub_stages/internal.py | 4 +- astrbot/core/provider/provider.py | 43 +++++++++++----- astrbot/dashboard/api/app.py | 21 +++++++- astrbot/dashboard/services/skills_service.py | 2 +- .../src/composables/useProviderSources.ts | 19 ++++++- dashboard/src/utils/providerMetadata.ts | 2 +- tests/agent/test_tool_loop_runner_history.py | 47 ++++++++++++++++++ tests/test_fastapi_v1_dashboard.py | 29 +++++++++++ tests/test_openai_source.py | 49 +++++++++++++++---- tests/unit/test_skills_service.py | 37 ++++++++++++++ 11 files changed, 255 insertions(+), 40 deletions(-) create mode 100644 tests/agent/test_tool_loop_runner_history.py create mode 100644 tests/unit/test_skills_service.py diff --git a/astrbot/core/agent/runners/tool_loop_agent_runner.py b/astrbot/core/agent/runners/tool_loop_agent_runner.py index 43f6beb99c..c1e4bccef6 100644 --- a/astrbot/core/agent/runners/tool_loop_agent_runner.py +++ b/astrbot/core/agent/runners/tool_loop_agent_runner.py @@ -194,7 +194,7 @@ async def _complete_with_assistant_response(self, llm_resp: LLMResponse) -> None parts.append(TextPart(text=llm_resp.completion_text)) if len(parts) == 0: logger.warning("LLM returned empty assistant message with no tool calls.") - self.run_context.messages.append(Message(role="assistant", content=parts)) + self._append_run_messages([Message(role="assistant", content=parts)]) try: await self.agent_hooks.on_agent_done(self.run_context, llm_resp) @@ -313,6 +313,7 @@ async def reset( ): m = await self._assemble_request_context_for_provider(request) messages.append(Message.model_validate(m)) + self._persistent_messages = list(messages) if request.system_prompt: messages.insert( 0, @@ -323,6 +324,23 @@ async def reset( self.stats = AgentStats() self.stats.start_time = time.time() + def _append_run_messages(self, messages: list[Message]) -> None: + """Append new agent messages to request and persistent histories. + + Args: + messages: New messages produced during the current agent run. + """ + self.run_context.messages.extend(messages) + self._persistent_messages.extend(messages) + + def get_persistent_messages(self) -> list[Message]: + """Return the complete conversation history for persistence. + + Returns: + A copy of the untrimmed messages accumulated during this agent run. + """ + return list(self._persistent_messages) + def _read_tool_hint(self) -> str: if self.read_tool is not None: return f"`{self.read_tool.name}`" @@ -926,9 +944,7 @@ async def step(self): tool_calls_result=tool_call_result_blocks, ) # record the assistant message with tool calls - self.run_context.messages.extend( - tool_calls_result.to_openai_messages_model() - ) + self._append_run_messages(tool_calls_result.to_openai_messages_model()) # If there are cached images and the model supports image input, # append a user message with images so LLM can see them @@ -960,8 +976,8 @@ async def step(self): ) ) if image_parts: - self.run_context.messages.append( - Message(role="user", content=image_parts) + self._append_run_messages( + [Message(role="user", content=image_parts)] ) logger.debug( f"Appended {len(cached_images)} cached image(s) to context for LLM review" @@ -988,11 +1004,13 @@ async def step_until_done( if self.req: self.req.func_tool = None # 注入提示词 - self.run_context.messages.append( - Message( - role="user", - content=self.MAX_STEPS_REACHED_PROMPT, - ) + self._append_run_messages( + [ + Message( + role="user", + content=self.MAX_STEPS_REACHED_PROMPT, + ) + ] ) # 再执行最后一步 async for resp in self.step(): @@ -1411,7 +1429,7 @@ async def _finalize_aborted_step( if llm_resp.completion_text: parts.append(TextPart(text=llm_resp.completion_text)) if parts: - self.run_context.messages.append(Message(role="assistant", content=parts)) + self._append_run_messages([Message(role="assistant", content=parts)]) try: await self.agent_hooks.on_agent_done(self.run_context, llm_resp) diff --git a/astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py b/astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py index 7b01dd2dc7..7346b0a294 100644 --- a/astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py +++ b/astrbot/core/pipeline/process_stage/method/agent_sub_stages/internal.py @@ -329,7 +329,7 @@ async def process( event, req, agent_runner.get_final_llm_resp(), - agent_runner.run_context.messages, + agent_runner.get_persistent_messages(), agent_runner.stats, user_aborted=agent_runner.was_aborted(), ) @@ -404,7 +404,7 @@ async def process( event, req, final_resp, - agent_runner.run_context.messages, + agent_runner.get_persistent_messages(), agent_runner.stats, user_aborted=agent_runner.was_aborted(), ) diff --git a/astrbot/core/provider/provider.py b/astrbot/core/provider/provider.py index 0cc9f1ca1c..3c083a241c 100644 --- a/astrbot/core/provider/provider.py +++ b/astrbot/core/provider/provider.py @@ -172,19 +172,38 @@ async def text_chat_stream( raise NotImplementedError() async def pop_record(self, context: list) -> None: - """弹出 context 第一条非系统提示词对话记录""" - poped = 0 - indexs_to_pop = [] - for idx, record in enumerate(context): - if record["role"] == "system": - continue - indexs_to_pop.append(idx) - poped += 1 - if poped == 2: - break + """Remove the oldest complete conversation turn from a request context. - for idx in reversed(indexs_to_pop): - context.pop(idx) + Args: + context: OpenAI-compatible message dictionaries to trim in place. + """ + first_message_index = next( + ( + index + for index, record in enumerate(context) + if record.get("role") != "system" + ), + None, + ) + if first_message_index is None: + return + + next_user_index = next( + ( + index + for index in range(first_message_index + 1, len(context)) + if context[index].get("role") == "user" + ), + len(context), + ) + context[:] = [ + record + for index, record in enumerate(context) + if not ( + first_message_index <= index < next_user_index + and record.get("role") != "system" + ) + ] def _ensure_message_to_dicts( self, diff --git a/astrbot/dashboard/api/app.py b/astrbot/dashboard/api/app.py index 00ceb11cfc..f5e6eca315 100644 --- a/astrbot/dashboard/api/app.py +++ b/astrbot/dashboard/api/app.py @@ -5,7 +5,7 @@ from fastapi import FastAPI, HTTPException, Request from fastapi.responses import JSONResponse -from astrbot.core import LogBroker +from astrbot.core import LogBroker, logger from astrbot.core.core_lifecycle import AstrBotCoreLifecycle from astrbot.core.db import BaseDatabase from astrbot.dashboard.responses import ApiError, error @@ -162,6 +162,25 @@ async def http_error_handler(_request: Request, exc: HTTPException): detail = exc.detail if isinstance(exc.detail, str) else "Request failed" return JSONResponse(error(detail), status_code=exc.status_code) + @app.exception_handler(Exception) + async def unhandled_error_handler(request: Request, exc: Exception): + """Log unexpected API failures and return a safe JSON response. + + Args: + request: Dashboard request that raised the exception. + exc: Unhandled exception raised while serving the request. + + Returns: + A generic JSON error response without internal exception details. + """ + logger.exception( + "Unhandled dashboard API exception for %s %s: %s", + request.method, + request.url.path, + exc, + ) + return JSONResponse(error("Internal server error"), status_code=500) + # Legacy dashboard routes keep old /api/* callers working without entering OpenAPI. app.include_router(legacy_api_keys_router) app.include_router(legacy_auth_router) diff --git a/astrbot/dashboard/services/skills_service.py b/astrbot/dashboard/services/skills_service.py index 4bcc6f8810..5a0f3bc45c 100644 --- a/astrbot/dashboard/services/skills_service.py +++ b/astrbot/dashboard/services/skills_service.py @@ -443,7 +443,7 @@ def prepare_skill_archive(self, name: str) -> SkillArchive: export_dir = Path(get_astrbot_temp_path()) / "skill_exports" export_dir.mkdir(parents=True, exist_ok=True) zip_base = export_dir / skill_name - zip_path = zip_base.with_suffix(".zip") + zip_path = Path(f"{zip_base}.zip") if zip_path.exists(): zip_path.unlink() diff --git a/dashboard/src/composables/useProviderSources.ts b/dashboard/src/composables/useProviderSources.ts index d0616f3b5b..b02d53f747 100644 --- a/dashboard/src/composables/useProviderSources.ts +++ b/dashboard/src/composables/useProviderSources.ts @@ -67,6 +67,7 @@ export function useProviderSources(options: UseProviderSourcesOptions) { const modelSearch = ref('') let suppressSourceWatch = false + let persistedProviderSourceIds = new Set() const providerTypes = computed(() => [ { value: 'chat_completion', label: tm('providers.tabs.chatCompletion'), icon: 'mdi-message-text' }, @@ -454,8 +455,14 @@ export function useProviderSources(options: UseProviderSourcesOptions) { ) if (!confirmed) return + const sourceId = String(source.id || '') + const isPersisted = persistedProviderSourceIds.has(sourceId) + try { - await providerApi.deleteSource(source.id) + if (isPersisted) { + await providerApi.deleteSource(sourceId) + persistedProviderSourceIds.delete(sourceId) + } providers.value = providers.value.filter((p) => p.provider_source_id !== source.id) providerSources.value = providerSources.value.filter((s) => s.id !== source.id) @@ -464,13 +471,18 @@ export function useProviderSources(options: UseProviderSourcesOptions) { selectedProviderSource.value = null selectedProviderSourceOriginalId.value = null editableProviderSource.value = null + availableModels.value = [] + modelMetadata.value = {} + isSourceModified.value = false } showMessage(tm('providerSources.deleteSuccess')) } catch (error: any) { showMessage(error.message || tm('providerSources.deleteError'), 'error') } finally { - await loadConfig() + if (isPersisted) { + await loadConfig() + } } } @@ -691,6 +703,9 @@ export function useProviderSources(options: UseProviderSourcesOptions) { providerTemplates.value = configSchema.value.provider.config_template } providerSources.value = response.data.data.provider_sources || [] + persistedProviderSourceIds = new Set( + providerSources.value.map((source: any) => String(source.id || '')) + ) modelMetadata.value = (response.data.data.model_metadata || {}) as Record providers.value = response.data.data.providers || [] } diff --git a/dashboard/src/utils/providerMetadata.ts b/dashboard/src/utils/providerMetadata.ts index 3d8faff4d4..0d5ce3c036 100644 --- a/dashboard/src/utils/providerMetadata.ts +++ b/dashboard/src/utils/providerMetadata.ts @@ -23,7 +23,7 @@ export function contextLimit( provider: ProviderMetadataSource | null | undefined, metadata?: ProviderModelMetadata | null ): number { - const context = Number(metadata?.limit?.context || provider?.max_context_tokens || 0) + const context = Number(provider?.max_context_tokens || metadata?.limit?.context || 0) return Number.isFinite(context) && context > 0 ? context : 0 } diff --git a/tests/agent/test_tool_loop_runner_history.py b/tests/agent/test_tool_loop_runner_history.py new file mode 100644 index 0000000000..9c01586909 --- /dev/null +++ b/tests/agent/test_tool_loop_runner_history.py @@ -0,0 +1,47 @@ +import pytest + +from astrbot.core.agent.hooks import BaseAgentRunHooks +from astrbot.core.agent.run_context import ContextWrapper +from astrbot.core.agent.runners.tool_loop_agent_runner import ToolLoopAgentRunner +from astrbot.core.agent.tool_executor import BaseFunctionToolExecutor +from astrbot.core.provider.entities import ProviderRequest +from astrbot.core.provider.sources.openai_source import ProviderOpenAIOfficial + + +@pytest.mark.asyncio +async def test_persistent_history_is_not_replaced_by_trimmed_request_context(): + provider = ProviderOpenAIOfficial( + provider_config={ + "id": "test-openai", + "type": "openai_chat_completion", + "model": "gpt-4o-mini", + "key": ["test-key"], + }, + provider_settings={}, + ) + request = ProviderRequest( + prompt="current request", + contexts=[ + {"role": "user", "content": "old request"}, + {"role": "assistant", "content": "old response"}, + ], + ) + try: + runner = ToolLoopAgentRunner() + await runner.reset( + provider=provider, + request=request, + run_context=ContextWrapper(context=None), + tool_executor=BaseFunctionToolExecutor(), + agent_hooks=BaseAgentRunHooks(), + ) + + runner.run_context.messages = runner.run_context.messages[-1:] + + assert [message.role for message in runner.get_persistent_messages()] == [ + "user", + "assistant", + "user", + ] + finally: + await provider.terminate() diff --git a/tests/test_fastapi_v1_dashboard.py b/tests/test_fastapi_v1_dashboard.py index f129c1e8de..9c8ccadca8 100644 --- a/tests/test_fastapi_v1_dashboard.py +++ b/tests/test_fastapi_v1_dashboard.py @@ -991,6 +991,35 @@ def _jwt_headers() -> dict[str, str]: return {"Authorization": f"Bearer {token}"} +@pytest.mark.asyncio +async def test_unhandled_dashboard_exception_returns_json_error( + asgi_app: FastAPI, + monkeypatch: pytest.MonkeyPatch, +): + async def raise_unhandled_error(): + raise RuntimeError("sensitive internal detail") + + monkeypatch.setattr( + asgi_app.state.services.stats, + "get_start_time", + raise_unhandled_error, + ) + transport = httpx.ASGITransport(app=asgi_app, raise_app_exceptions=False) + async with httpx.AsyncClient( + transport=transport, + base_url="http://testserver", + ) as client: + response = await client.get("/api/v1/stats/start-time") + + assert response.status_code == 500 + assert response.headers["content-type"].startswith("application/json") + assert response.json() == { + "status": "error", + "message": "Internal server error", + } + assert "sensitive internal detail" not in response.text + + @pytest.mark.asyncio async def test_public_versions_route_uses_static_folder( fake_core_lifecycle, diff --git a/tests/test_openai_source.py b/tests/test_openai_source.py index a45a232938..39109e994e 100644 --- a/tests/test_openai_source.py +++ b/tests/test_openai_source.py @@ -1,6 +1,7 @@ import base64 import builtins from io import BytesIO +from pathlib import Path from types import SimpleNamespace import httpx @@ -60,6 +61,36 @@ def _make_groq_provider(overrides: dict | None = None) -> ProviderGroq: ) +@pytest.mark.asyncio +async def test_pop_record_removes_complete_oldest_tool_call_turn(): + provider = ProviderOpenAIOfficial.__new__(ProviderOpenAIOfficial) + context = [ + {"role": "system", "content": "system"}, + {"role": "user", "content": "first request"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call-1", + "type": "function", + "function": {"name": "search", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call-1", "content": "result"}, + {"role": "assistant", "content": "first response"}, + {"role": "user", "content": "second request"}, + ] + + await provider.pop_record(context) + + assert context == [ + {"role": "system", "content": "system"}, + {"role": "user", "content": "second request"}, + ] + + def test_create_http_client_uses_openai_httpx_module(monkeypatch): captured: dict[str, object] = {} @@ -902,21 +933,21 @@ async def test_prepare_chat_payload_materializes_context_file_uri_image_urls(tmp await provider.terminate() -def test_file_uri_to_path_preserves_windows_drive_letter(): - assert file_uri_to_path("file:///C:/tmp/quoted-image.png") == ( - "C:/tmp/quoted-image.png" +def test_file_uri_to_path_normalizes_windows_drive_letter(): + assert file_uri_to_path("file:///C:/tmp/quoted-image.png") == str( + Path("C:/tmp/quoted-image.png") ) -def test_file_uri_to_path_preserves_windows_netloc_drive_letter(): - assert file_uri_to_path("file://C:/tmp/quoted-image.png") == ( - "C:/tmp/quoted-image.png" +def test_file_uri_to_path_normalizes_windows_netloc_drive_letter(): + assert file_uri_to_path("file://C:/tmp/quoted-image.png") == str( + Path("C:/tmp/quoted-image.png") ) -def test_file_uri_to_path_preserves_remote_netloc_as_unc_path(): - assert file_uri_to_path("file://server/share/quoted-image.png") == ( - "//server/share/quoted-image.png" +def test_file_uri_to_path_normalizes_remote_netloc_as_unc_path(): + assert file_uri_to_path("file://server/share/quoted-image.png") == str( + Path("//server/share/quoted-image.png") ) diff --git a/tests/unit/test_skills_service.py b/tests/unit/test_skills_service.py new file mode 100644 index 0000000000..aa46532e05 --- /dev/null +++ b/tests/unit/test_skills_service.py @@ -0,0 +1,37 @@ +from pathlib import Path + +import astrbot.dashboard.services.skills_service as skills_service_module +from astrbot.dashboard.services.skills_service import SkillsService + + +def test_prepare_skill_archive_preserves_versioned_skill_name( + monkeypatch, tmp_path: Path +): + skills_root = tmp_path / "skills" + skill_dir = skills_root / "skill-writing-1.0.0" + skill_dir.mkdir(parents=True) + (skill_dir / "SKILL.md").write_text("# Skill", encoding="utf-8") + + class FakeSkillManager: + @staticmethod + def is_sandbox_only_skill(_name: str) -> bool: + return False + + @staticmethod + def is_plugin_skill(_name: str) -> bool: + return False + + FakeSkillManager.skills_root = skills_root + + monkeypatch.setattr(skills_service_module, "SkillManager", FakeSkillManager) + monkeypatch.setattr( + skills_service_module, + "get_astrbot_temp_path", + lambda: str(tmp_path / "temp"), + ) + + archive = SkillsService(None).prepare_skill_archive("skill-writing-1.0.0") + + assert archive.path.name == "skill-writing-1.0.0.zip" + assert archive.path.is_file() + assert archive.filename == "skill-writing-1.0.0.zip"