Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
43 changes: 43 additions & 0 deletions test/query_agent/test_query_model.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,13 @@
import json
from contextlib import asynccontextmanager, contextmanager
from typing import Annotated, List

import httpx
import pytest
from httpx_sse import ServerSentEvent
from pydantic import ValidationError

from weaviate_agents.classes.media import GeneratedImage, GeneratedImageOptions
from weaviate_agents.classes.query import (
AskModeResponse,
ProgressMessage,
Expand Down Expand Up @@ -1430,6 +1432,47 @@ def fake_post_with_capture(url, headers=None, json=None, timeout=None):
assert captured["json"]["query"] == {"messages": chat_messages}


def test_ask_with_annotated_image_output_format(monkeypatch):
captured = {}
image = {"image_prompt": "a red shoe", "base64": "QUJD"}
response = {**FAKE_ASK_SUCCESS_JSON, "final_answer": json.dumps(image)}

def fake_post_with_capture(url, headers=None, json=None, timeout=None):
captured["json"] = json
return FakeResponse(200, response)

monkeypatch.setattr(httpx, "post", fake_post_with_capture)
dummy_client = DummyClient()
agent = QueryAgent(
dummy_client, ["test_collection"], agents_host="http://dummy-agent"
)
agent._connection = dummy_client
agent._headers = dummy_client.additional_headers

with pytest.warns(UserWarning):
result = agent.ask(
"draw a shoe",
output_format=Annotated[
GeneratedImage, GeneratedImageOptions(shape="square")
],
)

# the schema is sent with its shape, and the answer parses back into a GeneratedImage
assert captured["json"]["output_format"]["X-image-shape"] == "square"
assert result.final_answer_parsed == GeneratedImage(**image)


def test_ask_rejects_unsupported_output_format():
dummy_client = DummyClient()
agent = QueryAgent(
dummy_client, ["test_collection"], agents_host="http://dummy-agent"
)
agent._connection = dummy_client

with pytest.raises(TypeError):
agent.ask("draw shoes", output_format=List[GeneratedImage])


def test_ask_failure(monkeypatch):
monkeypatch.setattr(httpx, "post", fake_post_failure)
dummy_client = DummyClient()
Expand Down
6 changes: 6 additions & 0 deletions test/test_imports.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,8 @@ def test_class_exports():
DatePropertyFilter,
DependentOperationStep,
FilterAndOr,
GeneratedImage,
GeneratedImageOptions,
GeoPropertyFilter,
IntegerArrayPropertyFilter,
IntegerPropertyAggregation,
Expand Down Expand Up @@ -133,11 +135,13 @@ def test_class_exports():
DatePropertyAggregation,
DatePropertyFilter,
DependentOperationStep,
GeneratedImageOptions,
IntegerArrayPropertyFilter,
IntegerPropertyAggregation,
IntegerPropertyFilter,
IsNullPropertyFilter,
GeoPropertyFilter,
GeneratedImage,
NumericMetrics,
Operations,
OperationStep,
Expand Down Expand Up @@ -193,6 +197,8 @@ def test_class_exports():
"QueryWithCollection",
"Source",
"ChatMessage",
"GeneratedImage",
"GeneratedImageOptions",
"ComparisonOperator",
"IntegerPropertyFilter",
"TextPropertyFilter",
Expand Down
2 changes: 1 addition & 1 deletion uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

3 changes: 3 additions & 0 deletions weaviate_agents/classes/__init__.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
from .core import Usage
from .media import GeneratedImage, GeneratedImageOptions
from .personalization import (
Persona,
PersonaInteraction,
Expand Down Expand Up @@ -115,6 +116,8 @@
"IsNullPropertyFilter",
"SearchModeResponseBase",
"ChatMessage",
"GeneratedImage",
"GeneratedImageOptions",
"AskModeResponse",
"ResearchModeResponse",
"ModelUnitUsage",
Expand Down
31 changes: 31 additions & 0 deletions weaviate_agents/classes/media.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,31 @@
from dataclasses import dataclass
from typing import Optional

from pydantic import BaseModel, ConfigDict, GetJsonSchemaHandler
from pydantic.json_schema import JsonSchemaValue, SkipJsonSchema
from pydantic_core import core_schema
from typing_extensions import Literal

IMAGE_KEYWORD = "X-query-agent-image"

ImageShape = Literal["square", "landscape", "portrait"]


class GeneratedImage(BaseModel):
model_config = ConfigDict(json_schema_extra={IMAGE_KEYWORD: True})

image_prompt: str
base64: SkipJsonSchema[str] # hidden from the LLM's schema; server fills it


@dataclass(frozen=True)
class GeneratedImageOptions:
shape: Optional[ImageShape] = None # unset falls back to the backend default

def __get_pydantic_json_schema__(
self, core_schema_: core_schema.CoreSchema, handler: GetJsonSchemaHandler
) -> JsonSchemaValue:
schema = handler(core_schema_)
if self.shape is not None:
schema["X-image-shape"] = self.shape
return schema
136 changes: 130 additions & 6 deletions weaviate_agents/query/query_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
from abc import ABC, abstractmethod
from json import JSONDecodeError, loads
from typing import (
Annotated,
Any,
AsyncGenerator,
Coroutine,
Expand All @@ -11,16 +12,19 @@
Optional,
TypeVar,
Union,
get_args,
get_origin,
overload,
)

import httpx
from httpx_sse import ServerSentEvent, aconnect_sse, connect_sse
from pydantic import BaseModel
from pydantic import BaseModel, TypeAdapter
from typing_extensions import deprecated
from weaviate.client import WeaviateAsyncClient, WeaviateClient

from weaviate_agents.base import ClientType, _BaseAgent
from weaviate_agents.classes.media import IMAGE_KEYWORD, GeneratedImage
from weaviate_agents.query.classes import (
AskModeResponse,
ParsedAskModeResponse,
Expand Down Expand Up @@ -92,10 +96,18 @@ def _prepare_request_body(

if isinstance(output_format, type) and issubclass(output_format, BaseModel):
output_format_json = output_format.model_json_schema()
elif _is_annotated_image(output_format):
# the schema comes from the Annotated form so it keeps the GeneratedImageOptions
output_format_json = TypeAdapter(output_format).json_schema()
elif isinstance(output_format, dict):
output_format_json = output_format
else:
elif output_format is None:
output_format_json = None
else:
raise TypeError(
"output_format must be a BaseModel subclass, a dict JSON schema, "
f"GeneratedImage, or Annotated[GeneratedImage, GeneratedImageOptions(...)], got {output_format!r}"
)

output = {
"query": query_request,
Expand Down Expand Up @@ -824,7 +836,8 @@ def ask(
The LLM will conform to the output format specified.
The `final_answer_parsed` output field in the response will also be of the type specified.
When passing a `dict`, the dictionary must conform to the Draft 2020-12 JSON Schema specification.
Defaults to `str`.
To generate an image, include a :class:`~weaviate_agents.classes.media.GeneratedImage` field within a Pydantic model (as `output_format`).
The base64 of the image will be returned in the corresponding field of `final_answer_parsed` as string.

Returns:
An instance of :class:`~weaviate_agents.query.classes.response.AskModeResponse` (or [:class:`~weaviate_agents.query.classes.response.ParsedAskModeResponse`]
Expand Down Expand Up @@ -855,6 +868,21 @@ def ask(
>>> result = agent.ask("What contracts were signed by Jane Doe in 2024? What were they about?", output_format=AnswerWithSources)
>>> print(type(result.final_answer_parsed))
<class 'AnswerWithSources'>

>>> from weaviate_agents.classes.media import GeneratedImage
>>> from base64 import b64decode
>>> class AnswerWithImage(BaseModel):
... answer: str
... image: GeneratedImage = Field(
... description="An advertisement for the product.",
... )
>>>
>>> agent = QueryAgent(
... client=client,
... collections=["ECommerce"],
... )
>>> result = agent.ask("What is the best selling product?", output_format=AnswerWithImage)
>>> image_bytes = b64decode(result.final_answer_parsed.image.base64) # base64 encoded image
"""
request_body = self._prepare_request_body(
query=query,
Expand Down Expand Up @@ -1136,6 +1164,10 @@ def ask_stream(
Whilst streaming, the :class:`~weaviate_agents.query.classes.response.StreamedTokens` will return delta text
tokens on the final answer as it is being constructed as raw string tokens, not a JSON object.
When the final answer is complete, the :class:`~weaviate_agents.query.classes.response.ParsedAskModeResponse` will be returned if ``include_final_state`` is ``True``.
To generate an image, include a :class:`~weaviate_agents.classes.media.GeneratedImage` field within a Pydantic model (as `output_format`).
The base64 of the image will be returned in the corresponding field of `final_answer_parsed` as string.
If a image is being generated, the streamed tokens of the final answer will not include the `"base64"` key, it is added afterwards
and will be present in the final result.

Returns:
A generator of the response stream.
Expand All @@ -1153,7 +1185,7 @@ def ask_stream(
... collections=["FinancialContracts"],
... )
>>> for result in agent.ask_stream("What are the terms of the contract signed by John Smith in May 2025?"):
... if isinstance(result, AskModeResponse):
... if isinstance(result, AskModeResponse): # this will also find ParsedAskModeResponse
... result.display()
... elif isinstance(result, StreamedTokens):
... print(result.delta, end='', flush=True)
Expand All @@ -1178,12 +1210,32 @@ def ask_stream(
... "What contracts were signed by Jane Doe in 2024? What were they about?",
... output_format=AnswerWithSources
... ):
... if isinstance(result, AskModeResponse): # this will also find ParsedAskModeResponse
... if isinstance(result, AskModeResponse):
... result.display()
... elif isinstance(result, StreamedTokens):
... print(result.delta, end='', flush=True)
... elif isinstance(result, ProgressMessage):
... print(result.message)

>>> from weaviate_agents.classes.media import GeneratedImage
>>> from base64 import b64decode
>>> class AnswerWithImage(BaseModel):
... answer: str
... image: GeneratedImage = Field(
... description="An advertisement for the product.",
... )
>>>
>>> agent = QueryAgent(
... client=client,
... collections=["ECommerce"],
... )
>>> for result in agent.ask_stream("What is the best selling product?", output_format=AnswerWithImage):
... if isinstance(result, AskModeResponse): # the final result; `base64` is only populated here
... image_bytes = b64decode(result.final_answer_parsed.image.base64)
... elif isinstance(result, StreamedTokens):
... print(result.delta, end='', flush=True)
... elif isinstance(result, ProgressMessage):
... print(result.message)
"""
request_body = self._prepare_request_body(
query=query,
Expand Down Expand Up @@ -1678,7 +1730,8 @@ async def ask(
The LLM will conform to the output format specified.
The `final_answer_parsed` output field in the response will also be of the type specified.
When passing a `dict`, the dictionary must conform to the Draft 2020-12 JSON Schema specification.
Defaults to `str`.
To generate an image, include a :class:`~weaviate_agents.classes.media.GeneratedImage` field within a Pydantic model (as `output_format`).
The base64 of the image will be returned in the corresponding field of `final_answer_parsed` as string.

Returns:
An instance of :class:`~weaviate_agents.query.classes.response.AskModeResponse` (or [:class:`~weaviate_agents.query.classes.response.ParsedAskModeResponse`] if ``output_format`` is not ``None``) which contains the final answer, sources,
Expand Down Expand Up @@ -1709,6 +1762,21 @@ async def ask(
>>> result = await agent.ask("What contracts were signed by Jane Doe in 2024? What were they about?", output_format=AnswerWithSources)
>>> print(type(result.final_answer_parsed))
<class 'AnswerWithSources'>

>>> from weaviate_agents.classes.media import GeneratedImage
>>> from base64 import b64decode
>>> class AnswerWithImage(BaseModel):
... answer: str
... image: GeneratedImage = Field(
... description="An advertisement for the product.",
... )
>>>
>>> agent = AsyncQueryAgent(
... client=client,
... collections=["ECommerce"],
... )
>>> result = await agent.ask("What is the best selling product?", output_format=AnswerWithImage)
>>> image_bytes = b64decode(result.final_answer_parsed.image.base64) # base64 encoded image
"""
request_body = self._prepare_request_body(
query=query,
Expand Down Expand Up @@ -1994,6 +2062,10 @@ async def ask_stream(
Whilst streaming, the :class:`~weaviate_agents.query.classes.response.StreamedTokens` will return delta text
tokens on the final answer as it is being constructed as raw string tokens, not a JSON object.
When the final answer is complete, the :class:`~weaviate_agents.query.classes.response.ParsedAskModeResponse` will be returned if ``include_final_state`` is ``True``.
To generate an image, include a :class:`~weaviate_agents.classes.media.GeneratedImage` field within a Pydantic model (as `output_format`).
The base64 of the image will be returned in the corresponding field of `final_answer_parsed` as string.
If a image is being generated, the streamed tokens of the final answer will not include the `"base64"` key, it is added afterwards
and will be present in the final result.

Returns:
A generator of the response stream.
Expand Down Expand Up @@ -2042,6 +2114,26 @@ async def ask_stream(
... print(result.delta, end='', flush=True)
... elif isinstance(result, ProgressMessage):
... print(result.message)

>>> from weaviate_agents.classes.media import GeneratedImage
>>> from base64 import b64decode
>>> class AnswerWithImage(BaseModel):
... answer: str
... image: GeneratedImage = Field(
... description="An advertisement for the product.",
... )
>>>
>>> agent = AsyncQueryAgent(
... client=client,
... collections=["ECommerce"],
... )
>>> async for result in agent.ask_stream("What is the best selling product?", output_format=AnswerWithImage):
... if isinstance(result, AskModeResponse): # the final result; `base64` is only populated here
... image_bytes = b64decode(result.final_answer_parsed.image.base64)
... elif isinstance(result, StreamedTokens):
... print(result.delta, end='', flush=True)
... elif isinstance(result, ProgressMessage):
... print(result.message)
"""
request_body = self._prepare_request_body(
query=query,
Expand Down Expand Up @@ -2507,6 +2599,12 @@ def _parse_ask_result(
)
return ParsedAskModeResponse[BaseModel](**response)

elif _is_annotated_image(output_format):
response["final_answer_parsed"] = TypeAdapter(output_format).validate_json(
response["final_answer"]
)
return ParsedAskModeResponse[BaseModel](**response)

elif isinstance(output_format, dict):
try:
response["final_answer_parsed"] = loads(response["final_answer"])
Expand All @@ -2518,3 +2616,29 @@ def _parse_ask_result(
return ParsedAskModeResponse[dict](**response)

return AskModeResponse(**response)


def _is_annotated_image(output_format: Any) -> bool:
"""Whether output_format is an Annotated[GeneratedImage, ...] root, e.g. carrying GeneratedImageOptions."""
if get_origin(output_format) is not Annotated:
return False
base = get_args(output_format)[0]
return isinstance(base, type) and issubclass(base, GeneratedImage)


def _schema_contains_media(
schema: Union[dict[str, Any], list[Any], type[BaseModel]],
) -> bool:
"""Return True if the schema (or any nested node) requests generated media.

Used to increase timeouts when any nested node is tagged with the media keyword.
"""
if isinstance(schema, type) and issubclass(schema, BaseModel):
schema = schema.model_json_schema()
if isinstance(schema, dict):
if IMAGE_KEYWORD in schema:
return True
return any(_schema_contains_media(value) for value in schema.values())
if isinstance(schema, list):
return any(_schema_contains_media(item) for item in schema)
return False
Loading