This commit is contained in:
Javier Martinez
2026-07-24 12:56:49 +02:00
parent b3fca5591a
commit f482732a41
10 changed files with 85 additions and 67 deletions

View File

@@ -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,

View File

@@ -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,
)

View File

@@ -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,

View File

@@ -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]

View File

@@ -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)

View File

@@ -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

View File

@@ -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",

View File

@@ -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()

View File

@@ -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,
)

View File

@@ -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(