diff --git a/private_gpt/arq/tasks/chat/resume.py b/private_gpt/arq/tasks/chat/resume.py index ceabe9f1..896956d3 100644 --- a/private_gpt/arq/tasks/chat/resume.py +++ b/private_gpt/arq/tasks/chat/resume.py @@ -13,9 +13,12 @@ from private_gpt.arq.tasks.chat.settings import ( ) from private_gpt.components.engines.chat.async_chat_engine import AsyncChatEngine from private_gpt.components.tools.remote_execution import ToolExecutionResponse +from private_gpt.components.tools.tool_execution_outcome import ( + ToolExecutionError, + ToolExecutionFailure, +) from private_gpt.components.tools.tool_scheduler import ToolSchedulerFactory from private_gpt.di import get_global_injector -from private_gpt.events.models import TextBlock from private_gpt.server.chat.chat_service import ChatService from private_gpt.settings.settings import settings @@ -120,8 +123,12 @@ def _timeout_response( return ToolExecutionResponse( tool_name=tool_name, tool_id=tool_id, - result_content=[TextBlock(text=message)], - is_error=True, + outcome=ToolExecutionFailure( + error=ToolExecutionError( + code="tool_timeout", + message=message, + ) + ), tool_message=ChatMessage( role="tool", content=message, diff --git a/private_gpt/celery/tasks/tools/tool_run_task.py b/private_gpt/celery/tasks/tools/tool_run_task.py index c8ea2d71..9a66e373 100644 --- a/private_gpt/celery/tasks/tools/tool_run_task.py +++ b/private_gpt/celery/tasks/tools/tool_run_task.py @@ -17,9 +17,12 @@ from private_gpt.components.tools.remote_execution import ( execute_tool_request, resolve_tool_execution_interceptors, ) +from private_gpt.components.tools.tool_execution_outcome import ( + ToolExecutionError, + ToolExecutionFailure, +) from private_gpt.components.tools.tool_scheduler import ToolSchedulerFactory from private_gpt.di import get_global_injector -from private_gpt.events.models import TextBlock from private_gpt.settings.settings import settings logger = logging.getLogger(__name__) @@ -74,8 +77,12 @@ async def tool_run_task(*, request_data: dict[str, Any]) -> dict[str, Any]: response = ToolExecutionResponse( tool_name=request.tool_name, tool_id=request.tool_id, - result_content=[TextBlock(text=str(exc))], - is_error=True, + outcome=ToolExecutionFailure( + error=ToolExecutionError( + message=str(exc), + exception_type=type(exc).__name__, + ) + ), tool_message=request_error_message(request, str(exc)), ) else: @@ -86,7 +93,7 @@ async def tool_run_task(*, request_data: dict[str, Any]) -> dict[str, Any]: message_id, request.tool_id, request.tool_name, - response.is_error, + isinstance(response.outcome, ToolExecutionFailure), ) logger.debug( @@ -116,7 +123,7 @@ async def tool_run_task(*, request_data: dict[str, Any]) -> dict[str, Any]: correlation_id, message_id, request.tool_id, - response.is_error, + isinstance(response.outcome, ToolExecutionFailure), _result_fragment(response), ) return response.model_dump(mode="json") @@ -141,8 +148,12 @@ def _duplicate_execution_response( return ToolExecutionResponse( tool_name=request.tool_name, tool_id=request.tool_id, - result_content=[TextBlock(text=message)], - is_error=True, + outcome=ToolExecutionFailure( + error=ToolExecutionError( + code="duplicate_execution", + message=message, + ) + ), tool_message=request_error_message(request, message), ) @@ -168,7 +179,7 @@ async def _notify_completion( def _result_fragment(response: ToolExecutionResponse) -> str: serialized = json.dumps( - response.model_dump(mode="json")["result_content"], + response.model_dump(mode="json")["outcome"], ensure_ascii=False, default=str, ) diff --git a/private_gpt/components/tools/builders/web_fetch_builder.py b/private_gpt/components/tools/builders/web_fetch_builder.py index 5b9d61e5..a233e40e 100644 --- a/private_gpt/components/tools/builders/web_fetch_builder.py +++ b/private_gpt/components/tools/builders/web_fetch_builder.py @@ -4,15 +4,15 @@ from injector import inject, singleton from private_gpt.components.chat.models.chat_config_models import ToolSpec from private_gpt.components.llm.llm_component import LLMComponent +from private_gpt.components.tools.events.adapters import WebFetchEventAdapter from private_gpt.components.tools.remote_execution import build_rebuild_metadata from private_gpt.components.tools.tool_names import WEB_FETCH_TOOL_NAME from private_gpt.components.tools.tool_placeholders import WEB_FETCH_TOOL_FN from private_gpt.components.web.web_scraper_service import WebScraperService from private_gpt.di import get_global_injector -from private_gpt.components.tools.events.adapters import WebFetchEventAdapter from private_gpt.events.models import ( - WebFetchResultBlock, ResultContentBlockType, + WebFetchResultBlock, ) @@ -47,9 +47,16 @@ class WebFetchToolBuilder: if not result.markdown_content: from private_gpt.events.models import TextBlock - return [TextBlock(text="No content could be fetched from the provided URL.")] - return [WebFetchResultBlock.from_markdown(url=url, markdown=result.markdown_content)] + return [ + TextBlock(text="No content could be fetched from the provided URL.") + ] + + return [ + WebFetchResultBlock.from_markdown( + url=url, markdown=result.markdown_content + ) + ] return ToolSpec.from_defaults( name=name, diff --git a/private_gpt/components/tools/builders/web_search_builder.py b/private_gpt/components/tools/builders/web_search_builder.py index b4fb33f7..1e6a351e 100644 --- a/private_gpt/components/tools/builders/web_search_builder.py +++ b/private_gpt/components/tools/builders/web_search_builder.py @@ -1,23 +1,20 @@ -import asyncio from typing import Any, Literal, cast from injector import inject, singleton from private_gpt.components.chat.models.chat_config_models import ToolSpec -from private_gpt.components.chunk.models import Website from private_gpt.components.llm.llm_component import LLMComponent +from private_gpt.components.tools.events.adapters import WebSearchEventAdapter from private_gpt.components.tools.remote_execution import build_rebuild_metadata from private_gpt.components.tools.tool_names import WEB_SEARCH_TOOL_NAME from private_gpt.components.tools.tool_placeholders import WEB_SEARCH_TOOL_FN from private_gpt.components.tools.types import ToolValidationMode -from private_gpt.components.web.web_search.models import WebSearchResult from private_gpt.components.web.web_search.web_search_service import WebSearchService from private_gpt.di import get_global_injector -from private_gpt.components.tools.events.adapters import WebSearchEventAdapter from private_gpt.events.models import ( - from_tool_output, + ResultContentBlockType, + WebSearchResultBlock, ) -from private_gpt.events.models import WebSearchResultBlock, ResultContentBlockType @singleton @@ -50,22 +47,6 @@ class WebSearchToolBuilder: async def validate_search() -> None: await self.web_search_service.validate() - def _sync_format_results( - content: list[WebSearchResult], - ) -> list[ResultContentBlockType]: - if not content: - return [ - TextBlock( - text="No results found for the given query.", - ) - ] - - websites = [Website.from_website_result(res) for res in content] - return [ - *from_tool_output(websites), - *[TextBlock(text=str(result)) for result in content], - ] - async def run_tool(query: str) -> list[ResultContentBlockType]: if validate == ToolValidationMode.LAZY: # It is not validated because that would imply another call; @@ -75,6 +56,7 @@ class WebSearchToolBuilder: results = await self.web_search_service.search(query, model_id=model_id) if not results: from private_gpt.events.models import TextBlock + return [TextBlock(text="No results found for the given query.")] return [WebSearchResultBlock.from_web_search_result(r) for r in results] diff --git a/private_gpt/components/tools/events/adapters.py b/private_gpt/components/tools/events/adapters.py index 507779ca..2d6083c0 100644 --- a/private_gpt/components/tools/events/adapters.py +++ b/private_gpt/components/tools/events/adapters.py @@ -85,6 +85,8 @@ class ClientToolEventAdapter(ToolEventAdapter): class ServerToolEventAdapter(ToolEventAdapter): id_prefix = "srvtoolu" + _FALLBACK = None + def build_tool_use( self, *, tool_id: str, tool_name: str, tool_input: dict ) -> ToolUseBlock: @@ -100,8 +102,14 @@ class ServerToolEventAdapter(ToolEventAdapter): self, *, tool_use_id: str, outcome: ToolExecutionOutcome ) -> ToolResultBlock: if not self.public_tool_name: - return ClientToolEventAdapter.build_tool_result( - self, tool_use_id=tool_use_id, outcome=outcome + if isinstance(outcome, ToolExecutionFailure): + return ClientToolResultBlock( + tool_use_id=tool_use_id, + content=outcome.error.message, + is_error=True, + ) + return ClientToolResultBlock( + tool_use_id=tool_use_id, content=outcome.content ) return self._build_server_result(tool_use_id=tool_use_id, outcome=outcome) diff --git a/private_gpt/events/interceptors/filter_zylon_event_interceptor.py b/private_gpt/events/interceptors/filter_zylon_event_interceptor.py index 57462882..139dfe55 100644 --- a/private_gpt/events/interceptors/filter_zylon_event_interceptor.py +++ b/private_gpt/events/interceptors/filter_zylon_event_interceptor.py @@ -30,9 +30,7 @@ class FilterZylonEventInterceptor(BaseEventInterceptor): async def coro() -> AsyncGenerator[Event, None]: active_blocks: set[str] = set() async for event in gen: - new_event = event.for_response_mode( - self._response_mode - ) + new_event = event.for_response_mode(self._response_mode) if not new_event: continue diff --git a/private_gpt/events/models/__init__.py b/private_gpt/events/models/__init__.py index cb48f301..ca15ca2d 100644 --- a/private_gpt/events/models/__init__.py +++ b/private_gpt/events/models/__init__.py @@ -79,10 +79,10 @@ from private_gpt.events.models._tool_result_blocks import ( ContentBlockType, ServerToolResultBlock, ServerToolResultBlockType, - WebFetchToolResultBlock, - WebSearchToolResultBlock, TextEditorCodeExecutionToolResultBlock, ToolResultBlock, + WebFetchToolResultBlock, + WebSearchToolResultBlock, ) _types = [ @@ -166,11 +166,6 @@ __all__ = [ "CitationsDelta", "ClientToolResultBlock", "ClientToolUseBlock", - "WebFetchResultBlock", - "WebFetchToolResultBlock", - "WebSearchResultBlock", - "WebSearchToolResultError", - "WebSearchToolResultBlock", "CodeExecutionToolResultErrorBlock", "Container", "ContainerUploadBlock", @@ -223,6 +218,11 @@ __all__ = [ "ToolResultBlock", "ToolUseBlock", "Usage", + "WebFetchResultBlock", + "WebFetchToolResultBlock", + "WebSearchResultBlock", + "WebSearchToolResultBlock", + "WebSearchToolResultError", "from_tool_output", "serialize_datetime", "to_llama_index_blocks", diff --git a/tests/components/tools/test_server_tool_events.py b/tests/components/tools/test_server_tool_events.py index 46809389..ecb7141b 100644 --- a/tests/components/tools/test_server_tool_events.py +++ b/tests/components/tools/test_server_tool_events.py @@ -1,7 +1,5 @@ import pytest -import pytest - from private_gpt.components.chat.models.chat_config_models import ToolSpec from private_gpt.components.tools.events.adapters import ( BashCodeExecutionEventAdapter, @@ -19,7 +17,6 @@ from private_gpt.events.models import ( BashCodeExecutionToolResultBlock, ClientToolResultBlock, ClientToolUseBlock, - ServerToolResultBlock, ServerToolUseBlock, ToolResultBlock, ToolUseBlock, @@ -78,10 +75,12 @@ def test_client_tool_resolves_default_client_adapter() -> None: assert tool_id.startswith("tool_") -@pytest.mark.skip(reason="Internal server tools without an Anthropic-native name now use ToolUseBlock/ToolResultBlock " - "(client-style) for SDK Message.content compatibility. " - "Only Anthropic-native server tools (web_search, bash_code_execution, etc.) " - "use ServerToolUseBlock.") +@pytest.mark.skip( + reason="Internal server tools without an Anthropic-native name now use ToolUseBlock/ToolResultBlock " + "(client-style) for SDK Message.content compatibility. " + "Only Anthropic-native server tools (web_search, bash_code_execution, etc.) " + "use ServerToolUseBlock." +) def test_server_tool_resolves_default_server_adapter() -> None: tool = _tool() adapter = tool.resolve_event_adapter() diff --git a/tests/models/anthropic/registry.py b/tests/models/anthropic/registry.py index afafd23d..37c79121 100644 --- a/tests/models/anthropic/registry.py +++ b/tests/models/anthropic/registry.py @@ -6,6 +6,8 @@ from anthropic.types.raw_message_delta_event import Delta as SDKMessageDelta from private_gpt.events.models import ( AudioBlock, + BashCodeExecutionResultBlock, + BashCodeExecutionToolResultBlock, BinaryBlock, CitationsDelta, ContainerUploadBlock, @@ -36,6 +38,10 @@ from private_gpt.events.models import ( SourceDelta, TextBlock, TextDelta, + TextEditorCodeExecutionCreateResultBlock, + TextEditorCodeExecutionStrReplaceResultBlock, + TextEditorCodeExecutionToolResultBlock, + TextEditorCodeExecutionViewResultBlock, ThinkingBlock, ThinkingDelta, TLDRBlock, @@ -46,12 +52,6 @@ from private_gpt.events.models import ( WebFetchResultBlock, WebFetchToolResultBlock, WebSearchResultBlock, - BashCodeExecutionResultBlock, - BashCodeExecutionToolResultBlock, - TextEditorCodeExecutionViewResultBlock, - TextEditorCodeExecutionCreateResultBlock, - TextEditorCodeExecutionStrReplaceResultBlock, - TextEditorCodeExecutionToolResultBlock, WebSearchToolResultBlock, WebSearchToolResultError, ) diff --git a/tests/server/chat/anthropic/test_anthropic_client.py b/tests/server/chat/anthropic/test_anthropic_client.py index 760be82e..5dddcc8b 100644 --- a/tests/server/chat/anthropic/test_anthropic_client.py +++ b/tests/server/chat/anthropic/test_anthropic_client.py @@ -217,14 +217,14 @@ def validate_response_structure( if has_tools: tool_use_blocks = [ - b for b in response.content - if b.type in {"tool_use", "server_tool_use"} + b for b in response.content if b.type in {"tool_use", "server_tool_use"} ] assert len(tool_use_blocks) == 1 if has_internal_tools: tool_result_blocks = [ - b for b in response.content + b + for b in response.content if b.type in {"tool_result", "server_tool_result"} ] assert len(tool_result_blocks) == 1 @@ -342,7 +342,10 @@ def test_sync_chat_streaming( ): if event.content_block.type in {"tool_use", "server_tool_use"}: tool_use_count += 1 - elif event.content_block.type in {"tool_result", "server_tool_result"}: + elif event.content_block.type in { + "tool_result", + "server_tool_result", + }: tool_result_count += 1 validate_streaming_response( @@ -439,7 +442,10 @@ async def test_async_chat_streaming( ): if event.content_block.type in {"tool_use", "server_tool_use"}: tool_use_count += 1 - elif event.content_block.type in {"tool_result", "server_tool_result"}: + elif event.content_block.type in { + "tool_result", + "server_tool_result", + }: tool_result_count += 1 validate_streaming_response(