mirror of
https://github.com/VectifyAI/PageIndex.git
synced 2026-10-02 07:44:37 +08:00
feat: local chat — three protocol surfaces over the agent tools (v0.2.10)
Local mode gains managed document QA: an agent over the #393 local tool set, reachable through three wire protocols, each 1:1 with the backend and with no translation layer. - chat_completions(): standard chat.completions semantics on any OpenAI-compatible backend (openai-agents engine). Final answer only, cross-turn aggregated usage, streaming as text pieces or chunk dicts (the existing cloud signature, now implemented locally; model and max_turns are local-only additions). - responses(): the agentic surface — OpenAI Responses format, the tool process is standard output items, streaming forwards native events (tool outputs emitted as response.output_item.done, the way the platform streams its own server-side tools). Round-tripping output into the next input keeps provider prompt-cache prefix continuity and the agent's memory — live-verified: the follow-up call answered from round-tripped tool output with zero new tool calls. - messages(): Anthropic-native via the SDK's own tool runner (new pageindex[anthropic] extra, floor 0.68.0 verified for tool_runner/beta_tool(input_schema)). tool_use/tool_result round-trip is the format's native behavior; the envelope is the final message with aggregated usage plus the full new-turn sequence; the managed system blocks carry cache_control breakpoints. Shared skeleton: thin chat header + the local AGENT_INSTRUCTIONS (caller system content is appended, not rejected), the doc_id targeting block as a leading context item (factored out of build_agent_instructions), read-only toolset, structural-only validation (no arbitrary caps — backend limits govern), sampling params passed through, per-run tracing disabled, enable_citations rejected as cloud-only. Design basis is industry-standard formats rather than the cloud chat endpoint; responses()/messages() raise on cloud clients until the cloud converges. Tests run the real engines against scripted backends (a Model fake for openai-agents, a mock HTTP transport under the real anthropic SDK) with real tool execution against a seeded store, including the round-trip prefix-extension assertions on both engines.
This commit is contained in:
@@ -31,5 +31,5 @@ jobs:
|
||||
cache: pip
|
||||
- run: pip install -r requirements.txt pytest
|
||||
- if: matrix.agent-frameworks == 'with'
|
||||
run: pip install openai-agents claude-agent-sdk
|
||||
run: pip install openai-agents claude-agent-sdk anthropic
|
||||
- run: python -m pytest -q
|
||||
|
||||
+23
-17
@@ -1450,17 +1450,17 @@ def _base_instructions(client) -> str:
|
||||
return instructions
|
||||
|
||||
|
||||
def build_agent_instructions(client, doc_id=None) -> str:
|
||||
"""Orchestration guidance for document QA agents; with doc_id, appends
|
||||
the target documents and directs the agent to work within them. Raises
|
||||
when a doc_id's name is shadowed by a newer same-name document — the
|
||||
def doc_targeting_block(client, doc_id) -> Optional[str]:
|
||||
"""The doc_id targeting text: names, metadata, and the directive to work
|
||||
within those documents. Shared by agent_instructions and the local chat
|
||||
surfaces (which place it as a leading conversation item). Raises when a
|
||||
doc_id's name is shadowed by a newer same-name document — the
|
||||
name-addressed tools could not reach it."""
|
||||
base = _base_instructions(client)
|
||||
if doc_id is None:
|
||||
return base
|
||||
return None
|
||||
doc_ids = [doc_id] if isinstance(doc_id, str) else list(doc_id)
|
||||
if not doc_ids:
|
||||
return base
|
||||
return None
|
||||
details = [client.get_document(one_id) for one_id in doc_ids]
|
||||
documents = _all_documents(client)
|
||||
for one_id, detail in zip(doc_ids, details):
|
||||
@@ -1476,18 +1476,24 @@ def build_agent_instructions(client, doc_id=None) -> str:
|
||||
)
|
||||
context = json.dumps(details, ensure_ascii=False)
|
||||
if len(details) == 1:
|
||||
block = (
|
||||
return (
|
||||
f"The user has specified document: {details[0].get('name')}\n"
|
||||
f"Document metadata: {context}\n"
|
||||
"Use this document's name to retrieve its content with "
|
||||
"get_document_structure() and get_page_content()."
|
||||
)
|
||||
else:
|
||||
names = ", ".join(str(item.get("name")) for item in details)
|
||||
block = (
|
||||
f"The user has specified documents: {names}\n"
|
||||
f"Documents metadata: {context}\n"
|
||||
"Use these documents' names to retrieve their content with "
|
||||
"get_document_structure() and get_page_content()."
|
||||
)
|
||||
return base + "\n\n" + block
|
||||
names = ", ".join(str(item.get("name")) for item in details)
|
||||
return (
|
||||
f"The user has specified documents: {names}\n"
|
||||
f"Documents metadata: {context}\n"
|
||||
"Use these documents' names to retrieve their content with "
|
||||
"get_document_structure() and get_page_content()."
|
||||
)
|
||||
|
||||
|
||||
def build_agent_instructions(client, doc_id=None) -> str:
|
||||
"""Orchestration guidance for document QA agents; with doc_id, appends
|
||||
the target documents and directs the agent to work within them."""
|
||||
base = _base_instructions(client)
|
||||
block = doc_targeting_block(client, doc_id)
|
||||
return base if block is None else base + "\n\n" + block
|
||||
|
||||
+134
-11
@@ -344,38 +344,161 @@ class PageIndexClient:
|
||||
temperature: Optional[float] = None,
|
||||
stream_metadata: bool = False,
|
||||
enable_citations: bool = False,
|
||||
model: Optional[str] = None,
|
||||
max_turns: Optional[int] = None,
|
||||
) -> Union[dict[str, Any], Iterator[str], Iterator[dict[str, Any]]]:
|
||||
"""
|
||||
PageIndex Chat Completions, scoped to specific PageIndex documents.
|
||||
PageIndex Chat Completions: document QA in one call.
|
||||
|
||||
Cloud: the hosted chat endpoint. Local: a managed document-QA agent
|
||||
run over the local tools against your own LLM backend's
|
||||
/chat/completions (requires ``pageindex[openai]``; the OpenAI SDK's
|
||||
usual env config — OPENAI_API_KEY, OPENAI_BASE_URL — selects the
|
||||
backend, so any OpenAI-compatible server works). The response
|
||||
carries the final answer only; for the tool-use process and
|
||||
prompt-cache round-trip use ``responses()`` or ``messages()``.
|
||||
|
||||
Args:
|
||||
messages: Conversation messages with 'role' and 'content' keys.
|
||||
Local also accepts system/developer messages — their content
|
||||
is appended to the managed system prompt.
|
||||
stream: Enable streaming responses.
|
||||
doc_id: Document ID or list of IDs to scope the conversation.
|
||||
temperature: Sampling temperature (0.0-1.0).
|
||||
temperature: Sampling temperature, passed through to the model.
|
||||
stream_metadata: With stream=True, yield chunk dicts instead of
|
||||
text pieces.
|
||||
enable_citations: Enable citation instructions in responses.
|
||||
enable_citations: Cloud-only — local mode raises (citations need
|
||||
block-level OCR data local mode does not store).
|
||||
model: Local only — backend model name (defaults to
|
||||
``retrieve_model``). The cloud endpoint selects its own.
|
||||
max_turns: Local only — cap on agent turns per call.
|
||||
|
||||
Returns:
|
||||
- stream=False: complete response dict ({'id', 'object', 'created',
|
||||
'choices', 'usage'})
|
||||
- stream=True, stream_metadata=False: iterator of text chunks
|
||||
- stream=True, stream_metadata=True: iterator of chunk dicts
|
||||
|
||||
Local: not yet supported — raises PageIndexAPIError. Agent-based
|
||||
local chat arrives in a later release.
|
||||
"""
|
||||
return self._require_cloud(
|
||||
"chat_completions is not yet supported in local mode — it arrives "
|
||||
"in a later release. Create the client with an api_key to use "
|
||||
"cloud chat."
|
||||
).chat_completions(
|
||||
from .cloud_api import CloudAPI
|
||||
if not isinstance(self._api, CloudAPI):
|
||||
from .local_chat import run_chat_completions
|
||||
return run_chat_completions(
|
||||
self, messages, stream=stream, doc_id=doc_id,
|
||||
temperature=temperature, stream_metadata=stream_metadata,
|
||||
enable_citations=enable_citations, model=model,
|
||||
max_turns=max_turns,
|
||||
)
|
||||
if model is not None or max_turns is not None:
|
||||
raise PageIndexAPIError(
|
||||
"model and max_turns are local-mode parameters — the cloud "
|
||||
"chat endpoint selects its own model."
|
||||
)
|
||||
return self._api.chat_completions(
|
||||
messages=messages, stream=stream, doc_id=doc_id,
|
||||
temperature=temperature, stream_metadata=stream_metadata,
|
||||
enable_citations=enable_citations,
|
||||
)
|
||||
|
||||
def responses(
|
||||
self,
|
||||
input: Union[str, list[dict[str, Any]]],
|
||||
model: Optional[str] = None,
|
||||
stream: bool = False,
|
||||
doc_id: Optional[Union[str, list[str]]] = None,
|
||||
instructions: Optional[str] = None,
|
||||
temperature: Optional[float] = None,
|
||||
top_p: Optional[float] = None,
|
||||
max_turns: Optional[int] = None,
|
||||
) -> Union[dict[str, Any], Iterator[dict[str, Any]]]:
|
||||
"""
|
||||
Document QA over the OpenAI Responses protocol — the agentic surface.
|
||||
|
||||
Local only for now. Drives your backend's /responses end to end (no
|
||||
translation layer), so the ``output`` carries the whole process as
|
||||
standard items — messages, function calls, and function outputs
|
||||
(the SDK executes the tools). Append the returned ``output`` to your
|
||||
next call's ``input`` verbatim to keep provider prompt-cache prefix
|
||||
continuity and the agent's memory of what it already read.
|
||||
|
||||
Requires ``pageindex[openai]`` and a backend that supports the
|
||||
Responses API; backends that only speak chat.completions should use
|
||||
``chat_completions()``.
|
||||
|
||||
Args:
|
||||
input: A user message string, or a list of Responses input items
|
||||
(round-trip prior ``output`` items here).
|
||||
model: Backend model name (defaults to ``retrieve_model``).
|
||||
stream: Yield native Responses stream events as dicts; tool
|
||||
outputs are emitted as ``response.output_item.done`` events
|
||||
and the final event is ``response.completed``.
|
||||
doc_id: Document ID or list of IDs to scope the conversation.
|
||||
instructions: Appended to the managed system prompt.
|
||||
temperature / top_p: Passed through to the model.
|
||||
max_turns: Cap on agent turns per call.
|
||||
"""
|
||||
from .cloud_api import CloudAPI
|
||||
if isinstance(self._api, CloudAPI):
|
||||
raise PageIndexAPIError(
|
||||
"responses is not available on PageIndex cloud yet — it is "
|
||||
"a local-mode surface for now."
|
||||
)
|
||||
from .local_chat import run_responses
|
||||
return run_responses(
|
||||
self, input, model=model, stream=stream, doc_id=doc_id,
|
||||
instructions=instructions, temperature=temperature, top_p=top_p,
|
||||
max_turns=max_turns,
|
||||
)
|
||||
|
||||
def messages(
|
||||
self,
|
||||
messages: list[dict[str, Any]],
|
||||
model: str,
|
||||
max_tokens: int,
|
||||
stream: bool = False,
|
||||
doc_id: Optional[Union[str, list[str]]] = None,
|
||||
system: Optional[Union[str, list[dict[str, Any]]]] = None,
|
||||
temperature: Optional[float] = None,
|
||||
top_p: Optional[float] = None,
|
||||
top_k: Optional[int] = None,
|
||||
stop_sequences: Optional[list[str]] = None,
|
||||
max_turns: Optional[int] = None,
|
||||
) -> Union[dict[str, Any], Iterator[Any]]:
|
||||
"""
|
||||
Document QA over the Anthropic Messages protocol — Claude-native.
|
||||
|
||||
Local only for now. Drives Anthropic's /v1/messages via the
|
||||
Anthropic SDK's own tool runner (requires ``pageindex[anthropic]``;
|
||||
ANTHROPIC_API_KEY selects the backend). ``tool_use``/``tool_result``
|
||||
round-trip is the format's native behavior: the response is the
|
||||
final message envelope with cross-turn aggregated ``usage`` plus a
|
||||
``messages`` field — the full new turn sequence, valid for verbatim
|
||||
append to your history. The managed system prompt and the doc
|
||||
targeting block carry ``cache_control`` breakpoints.
|
||||
|
||||
Args:
|
||||
messages: Native Messages-format history (including prior
|
||||
tool_use/tool_result blocks on round-trip).
|
||||
model / max_tokens: Required by the Messages API; passed through.
|
||||
stream: Yield the native event stream across turns, verbatim.
|
||||
doc_id: Document ID or list of IDs to scope the conversation.
|
||||
system: Appended after the managed system blocks.
|
||||
temperature / top_p / top_k / stop_sequences: Passed through.
|
||||
max_turns: Cap on agent turns per call.
|
||||
"""
|
||||
from .cloud_api import CloudAPI
|
||||
if isinstance(self._api, CloudAPI):
|
||||
raise PageIndexAPIError(
|
||||
"messages is not available on PageIndex cloud yet — it is "
|
||||
"a local-mode surface for now."
|
||||
)
|
||||
from .local_chat import run_messages
|
||||
return run_messages(
|
||||
self, messages, model=model, max_tokens=max_tokens,
|
||||
stream=stream, doc_id=doc_id, system=system,
|
||||
temperature=temperature, top_p=top_p, top_k=top_k,
|
||||
stop_sequences=stop_sequences, max_turns=max_turns,
|
||||
)
|
||||
|
||||
# ---------- DOCUMENT MANAGEMENT ----------
|
||||
|
||||
def get_document(self, doc_id: str) -> dict[str, Any]:
|
||||
|
||||
@@ -0,0 +1,464 @@
|
||||
"""Managed local chat: document-QA agents over the local tools.
|
||||
|
||||
Three methods, three wire protocols, 1:1 with the backend and no translation
|
||||
layer: ``chat_completions`` drives the backend's /chat/completions (any
|
||||
OpenAI-compatible backend, final answer only), ``responses`` drives
|
||||
/responses (process items are standard output; round-trip them for provider
|
||||
prompt-cache continuation and agent memory), ``messages`` drives Anthropic's
|
||||
/v1/messages via the SDK's own tool runner (tool_use/tool_result round-trip
|
||||
is the format's native behavior).
|
||||
|
||||
Content passes through untouched — the caller's messages, the model's
|
||||
answers, tool outputs, finish/stop reasons. The SDK owns only gatekeeping
|
||||
(structural validation), table-setting (managed instructions, tools, doc
|
||||
targeting), tool execution, and billing (usage aggregation, envelope ids).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any, Iterator, Optional, Union
|
||||
|
||||
from .agent_tools import (AGENT_INSTRUCTIONS, _local_description,
|
||||
_local_schema, call_tool, doc_targeting_block,
|
||||
tool_names)
|
||||
from .errors import PageIndexAPIError
|
||||
|
||||
CHAT_HEADER = (
|
||||
"You are PageIndex by Vectify AI, a document-focused assistant. "
|
||||
"Be concise, never use emojis, and do not expose tool names."
|
||||
)
|
||||
|
||||
|
||||
# ── shared: prompt, doc targeting, validation, sync bridges ──
|
||||
|
||||
def _managed_instructions(extra_system: list[str]) -> str:
|
||||
return "\n\n".join([CHAT_HEADER, AGENT_INSTRUCTIONS, *extra_system])
|
||||
|
||||
|
||||
def _doc_block(client, doc_id) -> Optional[str]:
|
||||
if doc_id is None:
|
||||
return None
|
||||
doc_ids = [doc_id] if isinstance(doc_id, str) else list(doc_id)
|
||||
missing = []
|
||||
for one_id in doc_ids:
|
||||
try:
|
||||
client.get_document(one_id)
|
||||
except PageIndexAPIError:
|
||||
missing.append(str(one_id))
|
||||
if missing:
|
||||
raise PageIndexAPIError(
|
||||
"Documents not found or access denied: " + ", ".join(missing)
|
||||
)
|
||||
return doc_targeting_block(client, doc_id)
|
||||
|
||||
|
||||
def _system_text(content: Any) -> str:
|
||||
"""Text of a system/developer message: a string, or text parts joined."""
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
texts = [part.get("text") for part in content
|
||||
if isinstance(part, dict) and isinstance(part.get("text"), str)]
|
||||
if texts:
|
||||
return "\n".join(texts)
|
||||
raise PageIndexAPIError(
|
||||
"system message content must be a string or a list of text parts."
|
||||
)
|
||||
|
||||
|
||||
def _split_chat_messages(messages) -> "tuple[list[str], list[dict]]":
|
||||
"""Validate the chat_completions surface's messages: system/developer
|
||||
content joins the managed instructions; user/assistant history passes
|
||||
through. Tool-history round-trips belong to responses()/messages()."""
|
||||
if not isinstance(messages, list) or not messages:
|
||||
raise PageIndexAPIError("messages must be a non-empty list.")
|
||||
system_texts: list[str] = []
|
||||
history: list[dict] = []
|
||||
for message in messages:
|
||||
if not isinstance(message, dict) or "role" not in message:
|
||||
raise PageIndexAPIError(
|
||||
"Each message must be a dict with 'role' and 'content'.")
|
||||
role = message["role"]
|
||||
if role in ("system", "developer"):
|
||||
system_texts.append(_system_text(message.get("content")))
|
||||
elif role in ("user", "assistant"):
|
||||
content = message.get("content")
|
||||
if not isinstance(content, str):
|
||||
raise PageIndexAPIError(
|
||||
"chat_completions content must be a string; for "
|
||||
"structured items use responses() or messages()."
|
||||
)
|
||||
history.append({"role": role, "content": content})
|
||||
else:
|
||||
raise PageIndexAPIError(
|
||||
f"Unsupported role for chat_completions: {role!r}. Tool "
|
||||
"history round-trips belong to responses() or messages()."
|
||||
)
|
||||
if not history:
|
||||
raise PageIndexAPIError("messages must contain a user or assistant "
|
||||
"message.")
|
||||
return system_texts, history
|
||||
|
||||
|
||||
def _run_sync(coro):
|
||||
try:
|
||||
asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
return asyncio.run(coro)
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool:
|
||||
return pool.submit(asyncio.run, coro).result()
|
||||
|
||||
|
||||
_SENTINEL = object()
|
||||
|
||||
|
||||
def _stream_sync(agen_factory) -> Iterator[Any]:
|
||||
"""Drive an async generator from a background thread; yield synchronously."""
|
||||
items: "queue.Queue[Any]" = queue.Queue()
|
||||
|
||||
def pump():
|
||||
async def consume():
|
||||
async for item in agen_factory():
|
||||
items.put(item)
|
||||
|
||||
try:
|
||||
asyncio.run(consume())
|
||||
except BaseException as exc: # re-raised on the consumer thread
|
||||
items.put(exc)
|
||||
return
|
||||
items.put(_SENTINEL)
|
||||
|
||||
threading.Thread(target=pump, daemon=True).start()
|
||||
while True:
|
||||
item = items.get()
|
||||
if item is _SENTINEL:
|
||||
return
|
||||
if isinstance(item, BaseException):
|
||||
raise item
|
||||
yield item
|
||||
|
||||
|
||||
# ── OpenAI engine (chat_completions / responses) ──
|
||||
|
||||
def _require_openai_agents(method: str) -> None:
|
||||
try:
|
||||
import agents # noqa: F401
|
||||
except ImportError as exc:
|
||||
raise PageIndexAPIError(
|
||||
f"{method} in local mode requires the OpenAI Agents SDK — "
|
||||
"pip install openai-agents (or pip install 'pageindex[openai]')."
|
||||
) from exc
|
||||
|
||||
|
||||
def _openai_model(protocol: str, model_name: str):
|
||||
"""The backend protocol driver — the seam tests replace with a fake."""
|
||||
from openai import AsyncOpenAI
|
||||
if protocol == "chat":
|
||||
from agents.models.openai_chatcompletions import (
|
||||
OpenAIChatCompletionsModel)
|
||||
return OpenAIChatCompletionsModel(model_name, AsyncOpenAI())
|
||||
from agents.models.openai_responses import OpenAIResponsesModel
|
||||
return OpenAIResponsesModel(model_name, openai_client=AsyncOpenAI())
|
||||
|
||||
|
||||
def _openai_agent(client, protocol: str, model_name: str, instructions: str,
|
||||
temperature, top_p):
|
||||
from agents import Agent, ModelSettings
|
||||
from .integrations.openai_agents import build_openai_tools
|
||||
return Agent(
|
||||
name="PageIndex",
|
||||
instructions=instructions,
|
||||
tools=build_openai_tools(client),
|
||||
model=_openai_model(protocol, model_name),
|
||||
model_settings=ModelSettings(temperature=temperature, top_p=top_p),
|
||||
)
|
||||
|
||||
|
||||
def _run_kwargs(max_turns) -> dict:
|
||||
# Managed runs never export traces — the caller opted into document QA,
|
||||
# not telemetry.
|
||||
from agents import RunConfig
|
||||
kwargs: dict = {"run_config": RunConfig(tracing_disabled=True)}
|
||||
if max_turns is not None:
|
||||
kwargs["max_turns"] = max_turns
|
||||
return kwargs
|
||||
|
||||
|
||||
def _openai_usage(raw_responses) -> dict:
|
||||
prompt = sum(r.usage.input_tokens for r in raw_responses)
|
||||
completion = sum(r.usage.output_tokens for r in raw_responses)
|
||||
return {"prompt_tokens": prompt, "completion_tokens": completion,
|
||||
"total_tokens": prompt + completion}
|
||||
|
||||
|
||||
def run_chat_completions(client, messages, stream: bool = False,
|
||||
doc_id=None, temperature: Optional[float] = None,
|
||||
stream_metadata: bool = False,
|
||||
enable_citations: bool = False,
|
||||
model: Optional[str] = None,
|
||||
max_turns: Optional[int] = None,
|
||||
) -> Union[dict, Iterator[str], Iterator[dict]]:
|
||||
_require_openai_agents("chat_completions")
|
||||
if enable_citations:
|
||||
raise PageIndexAPIError(
|
||||
"enable_citations is cloud-only — citations need block-level OCR "
|
||||
"data that local mode does not store."
|
||||
)
|
||||
system_texts, history = _split_chat_messages(messages)
|
||||
block = _doc_block(client, doc_id)
|
||||
items = ([{"role": "user", "content": block}] if block else []) + history
|
||||
model_name = model or client.retrieve_model
|
||||
agent = _openai_agent(client, "chat", model_name,
|
||||
_managed_instructions(system_texts),
|
||||
temperature, None)
|
||||
from agents import Runner
|
||||
if not stream:
|
||||
result = _run_sync(
|
||||
Runner.run(agent, input=items, **_run_kwargs(max_turns)))
|
||||
return {
|
||||
"id": f"chatcmpl-{uuid.uuid4().hex}",
|
||||
"object": "chat.completion",
|
||||
"created": int(time.time()),
|
||||
"model": model_name,
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant",
|
||||
"content": result.final_output or ""},
|
||||
"finish_reason": "stop",
|
||||
}],
|
||||
"usage": _openai_usage(result.raw_responses),
|
||||
}
|
||||
|
||||
chat_id = f"chatcmpl-{uuid.uuid4().hex}"
|
||||
created = int(time.time())
|
||||
|
||||
def chunk(delta: dict, finish=None) -> dict:
|
||||
return {
|
||||
"id": chat_id, "object": "chat.completion.chunk",
|
||||
"created": created, "model": model_name,
|
||||
"choices": [{"index": 0, "delta": delta,
|
||||
"finish_reason": finish}],
|
||||
}
|
||||
|
||||
async def agen():
|
||||
from openai.types.responses import ResponseTextDeltaEvent
|
||||
streamed = Runner.run_streamed(agent, input=items,
|
||||
**_run_kwargs(max_turns))
|
||||
first = True
|
||||
async for event in streamed.stream_events():
|
||||
if (event.type == "raw_response_event"
|
||||
and isinstance(event.data, ResponseTextDeltaEvent)):
|
||||
if first:
|
||||
yield chunk({"role": "assistant", "content": ""})
|
||||
first = False
|
||||
yield chunk({"content": event.data.delta})
|
||||
yield chunk({}, finish="stop")
|
||||
yield {
|
||||
"id": chat_id, "object": "chat.completion.chunk",
|
||||
"created": created, "model": model_name, "choices": [],
|
||||
"usage": _openai_usage(streamed.raw_responses),
|
||||
}
|
||||
|
||||
if stream_metadata:
|
||||
return _stream_sync(agen)
|
||||
return (piece["choices"][0]["delta"]["content"]
|
||||
for piece in _stream_sync(agen)
|
||||
if piece.get("choices")
|
||||
and "content" in piece["choices"][0]["delta"]
|
||||
and piece["choices"][0]["delta"]["content"])
|
||||
|
||||
|
||||
def run_responses(client, input, model: Optional[str] = None,
|
||||
stream: bool = False, doc_id=None,
|
||||
instructions: Optional[str] = None,
|
||||
temperature: Optional[float] = None,
|
||||
top_p: Optional[float] = None,
|
||||
max_turns: Optional[int] = None,
|
||||
) -> Union[dict, Iterator[dict]]:
|
||||
_require_openai_agents("responses")
|
||||
if isinstance(input, str):
|
||||
items = [{"role": "user", "content": input}]
|
||||
elif (isinstance(input, list) and input
|
||||
and all(isinstance(item, dict) for item in input)):
|
||||
items = list(input)
|
||||
else:
|
||||
raise PageIndexAPIError("input must be a non-empty string or list "
|
||||
"of item dicts.")
|
||||
block = _doc_block(client, doc_id)
|
||||
if block:
|
||||
items = [{"role": "user", "content": block}] + items
|
||||
extra = [instructions] if instructions else []
|
||||
model_name = model or client.retrieve_model
|
||||
agent = _openai_agent(client, "responses", model_name,
|
||||
_managed_instructions(extra), temperature, top_p)
|
||||
from agents import Runner
|
||||
|
||||
def envelope(output: list, raw_responses) -> dict:
|
||||
usage = _openai_usage(raw_responses)
|
||||
return {
|
||||
"id": f"resp_{uuid.uuid4().hex}",
|
||||
"object": "response",
|
||||
"created_at": int(time.time()),
|
||||
"model": model_name,
|
||||
"status": "completed",
|
||||
"output": output,
|
||||
"usage": {"input_tokens": usage["prompt_tokens"],
|
||||
"output_tokens": usage["completion_tokens"],
|
||||
"total_tokens": usage["total_tokens"]},
|
||||
}
|
||||
|
||||
if not stream:
|
||||
result = _run_sync(
|
||||
Runner.run(agent, input=[dict(item) for item in items],
|
||||
**_run_kwargs(max_turns)))
|
||||
output = result.to_input_list()[len(items):]
|
||||
return envelope(output, result.raw_responses)
|
||||
|
||||
async def agen():
|
||||
streamed = Runner.run_streamed(agent,
|
||||
input=[dict(item) for item in items],
|
||||
**_run_kwargs(max_turns))
|
||||
async for event in streamed.stream_events():
|
||||
if event.type == "raw_response_event":
|
||||
yield event.data.model_dump(exclude_unset=True)
|
||||
elif (event.type == "run_item_stream_event"
|
||||
and event.item.type == "tool_call_output_item"):
|
||||
# We are the tool executor, so we emit the output item the
|
||||
# way the platform streams its own server-side tools.
|
||||
yield {"type": "response.output_item.done",
|
||||
"item": dict(event.item.to_input_item())}
|
||||
output = streamed.to_input_list()[len(items):]
|
||||
yield {"type": "response.completed",
|
||||
"response": envelope(output, streamed.raw_responses)}
|
||||
|
||||
return _stream_sync(agen)
|
||||
|
||||
|
||||
# ── Anthropic engine (messages) ──
|
||||
|
||||
def _require_anthropic() -> None:
|
||||
try:
|
||||
import anthropic # noqa: F401
|
||||
except ImportError as exc:
|
||||
raise PageIndexAPIError(
|
||||
"messages in local mode requires the Anthropic SDK — "
|
||||
"pip install anthropic (or pip install 'pageindex[anthropic]')."
|
||||
) from exc
|
||||
|
||||
|
||||
def _anthropic_client():
|
||||
"""The backend client — the seam tests replace with a fake transport."""
|
||||
import anthropic
|
||||
return anthropic.Anthropic()
|
||||
|
||||
|
||||
def _runnable_tools(client) -> list:
|
||||
from anthropic import beta_tool
|
||||
|
||||
def make(name: str):
|
||||
def _fn(**kwargs: Any) -> str:
|
||||
return call_tool(client, name, kwargs)[0]
|
||||
|
||||
_fn.__name__ = name
|
||||
return beta_tool(_fn, name=name, description=_local_description(name),
|
||||
input_schema=_local_schema(name))
|
||||
|
||||
return [make(name) for name in tool_names()]
|
||||
|
||||
|
||||
def _anthropic_system(extra_system, block: Optional[str]) -> list[dict]:
|
||||
"""System blocks with cache_control on the stable managed prefix; the
|
||||
doc block and caller system content follow as their own blocks."""
|
||||
blocks = [{"type": "text",
|
||||
"text": CHAT_HEADER + "\n\n" + AGENT_INSTRUCTIONS,
|
||||
"cache_control": {"type": "ephemeral"}}]
|
||||
if block:
|
||||
blocks.append({"type": "text", "text": block,
|
||||
"cache_control": {"type": "ephemeral"}})
|
||||
if extra_system is None:
|
||||
return blocks
|
||||
if isinstance(extra_system, str):
|
||||
return blocks + [{"type": "text", "text": extra_system}]
|
||||
if isinstance(extra_system, list):
|
||||
return blocks + list(extra_system)
|
||||
raise PageIndexAPIError("system must be a string or a list of blocks.")
|
||||
|
||||
|
||||
def _anthropic_usage(turns) -> dict:
|
||||
fields = ("input_tokens", "output_tokens",
|
||||
"cache_creation_input_tokens", "cache_read_input_tokens")
|
||||
totals = {field: 0 for field in fields}
|
||||
for turn in turns:
|
||||
for field in fields:
|
||||
value = getattr(turn.usage, field, None)
|
||||
if isinstance(value, int):
|
||||
totals[field] += value
|
||||
return totals
|
||||
|
||||
|
||||
def run_messages(client, messages, model: str, max_tokens: int,
|
||||
stream: bool = False, doc_id=None, system=None,
|
||||
temperature: Optional[float] = None,
|
||||
top_p: Optional[float] = None,
|
||||
top_k: Optional[int] = None,
|
||||
stop_sequences: Optional[list[str]] = None,
|
||||
max_turns: Optional[int] = None,
|
||||
) -> Union[dict, Iterator[Any]]:
|
||||
_require_anthropic()
|
||||
if not isinstance(messages, list) or not messages:
|
||||
raise PageIndexAPIError("messages must be a non-empty list.")
|
||||
block = _doc_block(client, doc_id)
|
||||
prepared = [dict(message) for message in messages]
|
||||
passthrough = {key: value for key, value in {
|
||||
"temperature": temperature, "top_p": top_p, "top_k": top_k,
|
||||
"stop_sequences": stop_sequences,
|
||||
}.items() if value is not None}
|
||||
runner = _anthropic_client().beta.messages.tool_runner(
|
||||
max_tokens=max_tokens,
|
||||
messages=prepared,
|
||||
model=model,
|
||||
tools=_runnable_tools(client),
|
||||
system=_anthropic_system(system, block),
|
||||
stream=stream,
|
||||
**({"max_iterations": max_turns} if max_turns is not None else {}),
|
||||
**passthrough,
|
||||
)
|
||||
|
||||
if stream:
|
||||
def events() -> Iterator[Any]:
|
||||
for turn_stream in runner:
|
||||
for event in turn_stream:
|
||||
yield event
|
||||
return events()
|
||||
|
||||
turns = [turn for turn in runner]
|
||||
if not turns:
|
||||
raise PageIndexAPIError("The model returned no response.")
|
||||
captured: dict = {}
|
||||
|
||||
def capture(params):
|
||||
captured.update(params)
|
||||
return params
|
||||
|
||||
runner.set_messages_params(capture)
|
||||
conversation = list(captured.get("messages") or [])
|
||||
final = turns[-1]
|
||||
envelope = final.model_dump(mode="json")
|
||||
envelope["usage"] = _anthropic_usage(turns)
|
||||
# The full turn sequence (assistant tool_use + user tool_result + final),
|
||||
# valid for verbatim append to the caller's history. The runner appends
|
||||
# intermediate turns to its params but not the final assistant message.
|
||||
new_messages = conversation[len(prepared):]
|
||||
if not new_messages or new_messages[-1].get("role") != "assistant":
|
||||
new_messages = new_messages + [{
|
||||
"role": "assistant",
|
||||
"content": [block.model_dump(mode="json")
|
||||
for block in final.content],
|
||||
}]
|
||||
envelope["messages"] = new_messages
|
||||
return envelope
|
||||
+5
-1
@@ -1,6 +1,6 @@
|
||||
[tool.poetry]
|
||||
name = "pageindex"
|
||||
version = "0.2.9"
|
||||
version = "0.2.10"
|
||||
description = "Python SDK for PageIndex — reasoning-based, vectorless document retrieval, cloud and local"
|
||||
readme = "README.md"
|
||||
license = "MIT"
|
||||
@@ -42,10 +42,14 @@ claude-agent-sdk = { version = ">=0.1.0", optional = true }
|
||||
# 0.8.0 offloads sync tools to a thread; older versions run them inline and
|
||||
# a blocking bridge call would freeze the agent event loop.
|
||||
openai-agents = { version = ">=0.8.0", optional = true }
|
||||
# messages() drives the SDK's beta tool runner; 0.68.0 is the first release
|
||||
# with tool_runner(stream/system/max_iterations) and beta_tool(input_schema).
|
||||
anthropic = { version = ">=0.68.0", optional = true }
|
||||
|
||||
[tool.poetry.extras]
|
||||
claude = ["claude-agent-sdk"]
|
||||
openai = ["openai-agents"]
|
||||
anthropic = ["anthropic"]
|
||||
|
||||
[tool.poetry.group.dev.dependencies]
|
||||
pytest = ">=7.0"
|
||||
|
||||
@@ -643,8 +643,12 @@ def test_retrieval_endpoints_cloud_only(local_client):
|
||||
local_client.get_retrieval("any")
|
||||
|
||||
|
||||
def test_chat_completions_cloud_only(local_client):
|
||||
with pytest.raises(PageIndexAPIError, match="not yet supported in local mode"):
|
||||
def test_chat_completions_local_needs_agents_extra(local_client, monkeypatch):
|
||||
"""Local chat is implemented (see test_local_chat.py); without the
|
||||
openai-agents extra it raises the actionable install error."""
|
||||
import sys
|
||||
monkeypatch.setitem(sys.modules, "agents", None)
|
||||
with pytest.raises(PageIndexAPIError, match="pageindex\\[openai\\]"):
|
||||
local_client.chat_completions(
|
||||
messages=[{"role": "user", "content": "q"}])
|
||||
|
||||
|
||||
@@ -0,0 +1,426 @@
|
||||
"""Local chat surfaces: three protocols over fake backends — no network,
|
||||
no LLM keys. Tool execution runs for real against a seeded local store."""
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
import pageindex.local_chat as local_chat
|
||||
from pageindex import (PageIndexAPIError, PageIndexCloudClient,
|
||||
PageIndexLocalClient)
|
||||
from pageindex.local_chat import CHAT_HEADER
|
||||
from pageindex.local_store import DocStore
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||
|
||||
|
||||
def seed_doc(storage_path, doc_id, name):
|
||||
pages = [{"page_index": 1, "markdown": "Page one text about apples"}]
|
||||
tree = [{"title": "Doc", "node_id": "0000", "start_index": 1,
|
||||
"end_index": 1, "summary": "root summary", "text": "ROOT"}]
|
||||
meta = {
|
||||
"id": doc_id, "name": name, "description": "A test document",
|
||||
"status": "completed", "createdAt": "2026-08-01T10:00:00.123000",
|
||||
"pageNum": 1, "folderId": None, "metadata": None, "mode": "standard",
|
||||
}
|
||||
DocStore(storage_path).save_document(doc_id, meta, tree, pages)
|
||||
return doc_id
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def store_path(tmp_path):
|
||||
return str(tmp_path / "store")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client(store_path):
|
||||
return PageIndexLocalClient(storage_path=str(store_path))
|
||||
|
||||
|
||||
# ── OpenAI engine fakes (chat_completions / responses) ──
|
||||
|
||||
agents = pytest.importorskip("agents")
|
||||
|
||||
|
||||
def _msg_item(text):
|
||||
from openai.types.responses import (ResponseOutputMessage,
|
||||
ResponseOutputText)
|
||||
return ResponseOutputMessage(
|
||||
id="msg_1", type="message", role="assistant", status="completed",
|
||||
content=[ResponseOutputText(type="output_text", text=text,
|
||||
annotations=[])])
|
||||
|
||||
|
||||
def _call_item(name, arguments, call_id="call_1"):
|
||||
from openai.types.responses import ResponseFunctionToolCall
|
||||
return ResponseFunctionToolCall(
|
||||
id="fc_1", type="function_call", call_id=call_id, name=name,
|
||||
arguments=json.dumps(arguments), status="completed")
|
||||
|
||||
|
||||
def _usage():
|
||||
from agents.usage import Usage
|
||||
return Usage(requests=1, input_tokens=10, output_tokens=5,
|
||||
total_tokens=15)
|
||||
|
||||
|
||||
from agents.models.interface import Model # noqa: E402
|
||||
|
||||
|
||||
class FakeModel(Model):
|
||||
"""Scripted backend: one list of output items per model turn."""
|
||||
|
||||
def __init__(self, turns):
|
||||
self.turns = list(turns)
|
||||
self.inputs = []
|
||||
self.instructions = []
|
||||
|
||||
def _record(self, system_instructions, input):
|
||||
self.instructions.append(system_instructions)
|
||||
items = input if isinstance(input, list) else [input]
|
||||
self.inputs.append(
|
||||
[dict(item) if isinstance(item, dict) else item
|
||||
for item in items])
|
||||
|
||||
async def get_response(self, system_instructions, input, model_settings,
|
||||
tools, output_schema, handoffs, tracing,
|
||||
**kwargs):
|
||||
from agents.items import ModelResponse
|
||||
self._record(system_instructions, input)
|
||||
return ModelResponse(output=self.turns.pop(0), usage=_usage(),
|
||||
response_id=None)
|
||||
|
||||
async def stream_response(self, system_instructions, input,
|
||||
model_settings, tools, output_schema, handoffs,
|
||||
tracing, **kwargs):
|
||||
from openai.types.responses import (Response, ResponseCompletedEvent,
|
||||
ResponseTextDeltaEvent)
|
||||
from openai.types.responses.response_usage import (
|
||||
InputTokensDetails, OutputTokensDetails, ResponseUsage)
|
||||
self._record(system_instructions, input)
|
||||
output = self.turns.pop(0)
|
||||
sequence = 0
|
||||
for item in output:
|
||||
if item.type == "message":
|
||||
for piece in ("The ", "answer"):
|
||||
sequence += 1
|
||||
yield ResponseTextDeltaEvent(
|
||||
type="response.output_text.delta", delta=piece,
|
||||
content_index=0, item_id=item.id, output_index=0,
|
||||
logprobs=[], sequence_number=sequence)
|
||||
sequence += 1
|
||||
yield ResponseCompletedEvent(
|
||||
type="response.completed", sequence_number=sequence,
|
||||
response=Response(
|
||||
id="resp_fake", created_at=0.0, model="fake",
|
||||
object="response", output=output, parallel_tool_calls=False,
|
||||
tool_choice="auto", tools=[],
|
||||
usage=ResponseUsage(
|
||||
input_tokens=10, output_tokens=5, total_tokens=15,
|
||||
input_tokens_details=InputTokensDetails(
|
||||
cached_tokens=0, cache_write_tokens=0),
|
||||
output_tokens_details=OutputTokensDetails(
|
||||
reasoning_tokens=0))))
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fake_model(monkeypatch):
|
||||
state = {}
|
||||
|
||||
def install(turns):
|
||||
fake = FakeModel(turns)
|
||||
state["protocols"] = []
|
||||
|
||||
def factory(protocol, model_name):
|
||||
state["protocols"].append((protocol, model_name))
|
||||
return fake
|
||||
|
||||
monkeypatch.setattr(local_chat, "_openai_model", factory)
|
||||
return fake
|
||||
|
||||
install.state = state
|
||||
return install
|
||||
|
||||
|
||||
# ── chat_completions ──
|
||||
|
||||
def test_chat_completions_end_to_end(client, store_path, fake_model):
|
||||
seed_doc(store_path, "pi-a", "report.pdf")
|
||||
fake = fake_model([
|
||||
[_call_item("get_document", {"doc_name": "report.pdf"})],
|
||||
[_msg_item("The answer")],
|
||||
])
|
||||
result = client.chat_completions(
|
||||
[{"role": "user", "content": "What status?"}])
|
||||
assert result["id"].startswith("chatcmpl-")
|
||||
assert result["object"] == "chat.completion"
|
||||
assert result["choices"][0]["message"] == {"role": "assistant",
|
||||
"content": "The answer"}
|
||||
assert result["choices"][0]["finish_reason"] == "stop"
|
||||
assert result["usage"] == {"prompt_tokens": 20, "completion_tokens": 10,
|
||||
"total_tokens": 30}
|
||||
assert fake_model.state["protocols"][0][0] == "chat"
|
||||
# The tool ran for real: turn 2's input carries its output.
|
||||
turn2 = json.dumps(fake.inputs[1])
|
||||
assert "report.pdf" in turn2 and "completed" in turn2
|
||||
# Managed instructions: header + the local agent guidance.
|
||||
assert fake.instructions[0].startswith(CHAT_HEADER)
|
||||
assert "READING WORKFLOW" in fake.instructions[0]
|
||||
|
||||
|
||||
def test_chat_completions_system_and_doc_block(client, store_path, fake_model):
|
||||
doc_id = seed_doc(store_path, "pi-a", "report.pdf")
|
||||
fake = fake_model([[_msg_item("ok")]])
|
||||
client.chat_completions(
|
||||
[{"role": "system", "content": "Answer in French."},
|
||||
{"role": "user", "content": "hi"}],
|
||||
doc_id=doc_id)
|
||||
assert fake.instructions[0].endswith("Answer in French.")
|
||||
first_item = fake.inputs[0][0]
|
||||
assert "The user has specified document: report.pdf" in first_item["content"]
|
||||
|
||||
|
||||
def test_chat_completions_validation(client, store_path, fake_model):
|
||||
fake_model([[_msg_item("ok")]])
|
||||
with pytest.raises(PageIndexAPIError, match="cloud-only"):
|
||||
client.chat_completions([{"role": "user", "content": "x"}],
|
||||
enable_citations=True)
|
||||
with pytest.raises(PageIndexAPIError, match="responses\\(\\) or messages"):
|
||||
client.chat_completions([{"role": "tool", "content": "x"}])
|
||||
with pytest.raises(PageIndexAPIError, match="must be a string"):
|
||||
client.chat_completions([{"role": "user", "content": [1]}])
|
||||
with pytest.raises(PageIndexAPIError, match="non-empty"):
|
||||
client.chat_completions([])
|
||||
with pytest.raises(PageIndexAPIError,
|
||||
match="Documents not found or access denied: a, b"):
|
||||
client.chat_completions([{"role": "user", "content": "x"}],
|
||||
doc_id=["a", "b"])
|
||||
|
||||
|
||||
def test_chat_completions_stream_modes(client, store_path, fake_model):
|
||||
fake_model([[_msg_item("The answer")]])
|
||||
pieces = list(client.chat_completions(
|
||||
[{"role": "user", "content": "q"}], stream=True))
|
||||
assert pieces == ["The ", "answer"]
|
||||
|
||||
fake_model([[_msg_item("The answer")]])
|
||||
chunks = list(client.chat_completions(
|
||||
[{"role": "user", "content": "q"}], stream=True,
|
||||
stream_metadata=True))
|
||||
assert chunks[0]["choices"][0]["delta"] == {"role": "assistant",
|
||||
"content": ""}
|
||||
assert chunks[-2]["choices"][0]["finish_reason"] == "stop"
|
||||
assert chunks[-1]["choices"] == []
|
||||
assert chunks[-1]["usage"]["total_tokens"] == 15
|
||||
assert all(c["object"] == "chat.completion.chunk" for c in chunks[:-1])
|
||||
|
||||
|
||||
def test_chat_completions_missing_framework(client, monkeypatch):
|
||||
monkeypatch.setitem(sys.modules, "agents", None)
|
||||
with pytest.raises(PageIndexAPIError, match="pageindex\\[openai\\]"):
|
||||
client.chat_completions([{"role": "user", "content": "x"}])
|
||||
|
||||
|
||||
def test_cloud_guards():
|
||||
cloud = PageIndexCloudClient(api_key="pi-test-key")
|
||||
with pytest.raises(PageIndexAPIError, match="local-mode parameters"):
|
||||
cloud.chat_completions([{"role": "user", "content": "x"}], model="m")
|
||||
with pytest.raises(PageIndexAPIError, match="not available on PageIndex "
|
||||
"cloud yet"):
|
||||
cloud.responses("x")
|
||||
with pytest.raises(PageIndexAPIError, match="not available on PageIndex "
|
||||
"cloud yet"):
|
||||
cloud.messages([{"role": "user", "content": "x"}], model="m",
|
||||
max_tokens=10)
|
||||
|
||||
|
||||
# ── responses ──
|
||||
|
||||
def test_responses_end_to_end(client, store_path, fake_model):
|
||||
seed_doc(store_path, "pi-a", "report.pdf")
|
||||
fake = fake_model([
|
||||
[_call_item("get_document", {"doc_name": "report.pdf"})],
|
||||
[_msg_item("The answer")],
|
||||
])
|
||||
result = client.responses("What status?")
|
||||
assert result["id"].startswith("resp_")
|
||||
assert result["object"] == "response"
|
||||
assert result["status"] == "completed"
|
||||
assert result["usage"] == {"input_tokens": 20, "output_tokens": 10,
|
||||
"total_tokens": 30}
|
||||
assert fake_model.state["protocols"][0][0] == "responses"
|
||||
types = [item.get("type", "message") for item in result["output"]]
|
||||
assert "function_call" in types and "function_call_output" in types
|
||||
# The final item is the assistant answer.
|
||||
assert "The answer" in json.dumps(result["output"][-1])
|
||||
|
||||
|
||||
def test_responses_round_trip_extends_prefix(client, store_path, fake_model):
|
||||
"""The cache contract: a round-tripped call's first model input must
|
||||
extend the previous call's final model input item-for-item."""
|
||||
seed_doc(store_path, "pi-a", "report.pdf")
|
||||
first = fake_model([
|
||||
[_call_item("get_document", {"doc_name": "report.pdf"})],
|
||||
[_msg_item("The answer")],
|
||||
])
|
||||
result = client.responses("What status?")
|
||||
|
||||
second = fake_model([[_msg_item("Done")]])
|
||||
follow_up = ([{"role": "user", "content": "What status?"}]
|
||||
+ result["output"]
|
||||
+ [{"role": "user", "content": "and now?"}])
|
||||
client.responses(follow_up)
|
||||
previous_final = first.inputs[-1]
|
||||
assert second.inputs[0][:len(previous_final)] == previous_final
|
||||
|
||||
|
||||
def test_responses_stream_passthrough(client, store_path, fake_model):
|
||||
seed_doc(store_path, "pi-a", "report.pdf")
|
||||
fake_model([
|
||||
[_call_item("get_document", {"doc_name": "report.pdf"})],
|
||||
[_msg_item("The answer")],
|
||||
])
|
||||
events = list(client.responses("q", stream=True))
|
||||
types = [event.get("type") for event in events]
|
||||
assert "response.output_text.delta" in types
|
||||
tool_events = [event for event in events
|
||||
if event.get("type") == "response.output_item.done"
|
||||
and event.get("item", {}).get("type")
|
||||
== "function_call_output"]
|
||||
assert tool_events, types
|
||||
assert types[-1] == "response.completed"
|
||||
final = events[-1]["response"]
|
||||
assert final["status"] == "completed"
|
||||
assert final["usage"]["total_tokens"] == 30
|
||||
|
||||
|
||||
# ── messages (Anthropic engine) ──
|
||||
|
||||
anthropic = pytest.importorskip("anthropic")
|
||||
import httpx # noqa: E402 (anthropic depends on httpx)
|
||||
|
||||
|
||||
def _anthropic_message(content, stop_reason):
|
||||
return {
|
||||
"id": "msg_fake", "type": "message", "role": "assistant",
|
||||
"model": "claude-test", "content": content,
|
||||
"stop_reason": stop_reason, "stop_sequence": None,
|
||||
"usage": {"input_tokens": 10, "output_tokens": 5},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fake_anthropic(monkeypatch):
|
||||
state = {"calls": []}
|
||||
|
||||
def install(responses):
|
||||
state["calls"].clear()
|
||||
|
||||
def handler(request):
|
||||
state["calls"].append(json.loads(request.content))
|
||||
body = responses[len(state["calls"]) - 1]
|
||||
if isinstance(body, str): # pre-rendered SSE
|
||||
return httpx.Response(
|
||||
200, content=body.encode(),
|
||||
headers={"content-type": "text/event-stream"})
|
||||
return httpx.Response(200, json=body)
|
||||
|
||||
fake = anthropic.Anthropic(
|
||||
api_key="test",
|
||||
http_client=httpx.Client(transport=httpx.MockTransport(handler)))
|
||||
monkeypatch.setattr(local_chat, "_anthropic_client", lambda: fake)
|
||||
return state["calls"]
|
||||
|
||||
return install
|
||||
|
||||
|
||||
def test_messages_end_to_end(client, store_path, fake_anthropic):
|
||||
seed_doc(store_path, "pi-a", "report.pdf")
|
||||
calls = fake_anthropic([
|
||||
_anthropic_message(
|
||||
[{"type": "tool_use", "id": "tu_1", "name": "get_document",
|
||||
"input": {"doc_name": "report.pdf"}}], "tool_use"),
|
||||
_anthropic_message([{"type": "text", "text": "The answer"}],
|
||||
"end_turn"),
|
||||
])
|
||||
result = client.messages([{"role": "user", "content": "What status?"}],
|
||||
model="claude-test", max_tokens=100)
|
||||
assert result["stop_reason"] == "end_turn"
|
||||
assert result["content"][0]["text"] == "The answer"
|
||||
assert result["usage"]["input_tokens"] == 20
|
||||
assert result["usage"]["output_tokens"] == 10
|
||||
# Full new-turn sequence, valid for verbatim history append.
|
||||
roles = [message["role"] for message in result["messages"]]
|
||||
assert roles == ["assistant", "user", "assistant"]
|
||||
tool_result = json.dumps(result["messages"][1])
|
||||
assert "tool_result" in tool_result and "report.pdf" in tool_result
|
||||
|
||||
request = calls[0]
|
||||
assert request["system"][0]["text"].startswith(CHAT_HEADER)
|
||||
assert request["system"][0]["cache_control"] == {"type": "ephemeral"}
|
||||
browse = next(t for t in request["tools"]
|
||||
if t["name"] == "browse_documents")
|
||||
assert "folder_id" not in browse["input_schema"]["properties"]
|
||||
# Native prefix continuation: request 2 extends request 1's messages.
|
||||
assert calls[1]["messages"][:len(calls[0]["messages"])] \
|
||||
== calls[0]["messages"]
|
||||
|
||||
|
||||
def test_messages_doc_block_and_system(client, store_path, fake_anthropic):
|
||||
doc_id = seed_doc(store_path, "pi-a", "report.pdf")
|
||||
calls = fake_anthropic([
|
||||
_anthropic_message([{"type": "text", "text": "ok"}], "end_turn"),
|
||||
])
|
||||
client.messages([{"role": "user", "content": "hi"}], model="claude-test",
|
||||
max_tokens=100, doc_id=doc_id, system="Answer in French.")
|
||||
system = calls[0]["system"]
|
||||
assert "The user has specified document: report.pdf" in system[1]["text"]
|
||||
assert system[-1]["text"] == "Answer in French."
|
||||
|
||||
|
||||
def test_messages_stream_passthrough(client, store_path, fake_anthropic):
|
||||
sse = "\n".join([
|
||||
'event: message_start',
|
||||
'data: {"type":"message_start","message":{"id":"msg_1","type":"message","role":"assistant","model":"claude-test","content":[],"stop_reason":null,"stop_sequence":null,"usage":{"input_tokens":10,"output_tokens":1}}}',
|
||||
"",
|
||||
'event: content_block_start',
|
||||
'data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}',
|
||||
"",
|
||||
'event: content_block_delta',
|
||||
'data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"The answer"}}',
|
||||
"",
|
||||
'event: content_block_stop',
|
||||
'data: {"type":"content_block_stop","index":0}',
|
||||
"",
|
||||
'event: message_delta',
|
||||
'data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{"output_tokens":5}}',
|
||||
"",
|
||||
'event: message_stop',
|
||||
'data: {"type":"message_stop"}',
|
||||
"",
|
||||
"",
|
||||
])
|
||||
fake_anthropic([sse])
|
||||
events = list(client.messages([{"role": "user", "content": "q"}],
|
||||
model="claude-test", max_tokens=100,
|
||||
stream=True))
|
||||
types = [event.type for event in events]
|
||||
assert "content_block_delta" in types and "message_stop" in types
|
||||
|
||||
|
||||
def test_messages_validation(client, fake_anthropic):
|
||||
fake_anthropic([])
|
||||
with pytest.raises(PageIndexAPIError, match="non-empty"):
|
||||
client.messages([], model="claude-test", max_tokens=100)
|
||||
with pytest.raises(PageIndexAPIError,
|
||||
match="Documents not found or access denied"):
|
||||
client.messages([{"role": "user", "content": "x"}],
|
||||
model="claude-test", max_tokens=100, doc_id="ghost")
|
||||
|
||||
|
||||
def test_messages_missing_framework(client, monkeypatch):
|
||||
monkeypatch.setitem(sys.modules, "anthropic", None)
|
||||
with pytest.raises(PageIndexAPIError, match="pageindex\\[anthropic\\]"):
|
||||
client.messages([{"role": "user", "content": "x"}],
|
||||
model="claude-test", max_tokens=100)
|
||||
Reference in New Issue
Block a user