mirror of
https://github.com/imartinez/privateGPT.git
synced 2026-08-09 08:03:42 +00:00
fix: ty
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user