Skip to content

Commit c9ac5cd

Browse files
committed
Implement structured output support delegated to providers
1 parent f914c09 commit c9ac5cd

8 files changed

Lines changed: 381 additions & 7 deletions

File tree

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,48 @@
1+
import asyncio
2+
import os
3+
4+
import pydantic
5+
6+
import vercel_ai_sdk as ai
7+
8+
9+
class WeatherForecast(pydantic.BaseModel):
10+
city: str
11+
temperature: float
12+
conditions: str
13+
humidity: int
14+
wind_speed: float
15+
16+
17+
async def main() -> None:
18+
llm = ai.openai.OpenAIModel(
19+
model="anthropic/claude-sonnet-4",
20+
base_url="https://ai-gateway.vercel.sh/v1",
21+
api_key=os.environ.get("AI_GATEWAY_API_KEY"),
22+
)
23+
24+
messages = ai.make_messages(
25+
system="You are a weather assistant. Respond with realistic weather data.",
26+
user="What's the weather like in San Francisco right now?",
27+
)
28+
29+
# Streaming: watch the JSON arrive incrementally, get validated output at the end
30+
print("--- Streaming ---")
31+
async for msg in llm.stream(messages, output_type=WeatherForecast):
32+
if msg.text_delta:
33+
print(msg.text_delta, end="", flush=True)
34+
if msg.output:
35+
print(f"\n\nParsed: {msg.output}")
36+
37+
# Non-streaming: get the validated output directly
38+
print("\n--- Buffer ---")
39+
msg = await llm.buffer(messages, output_type=WeatherForecast)
40+
print(f"City: {msg.output.city}")
41+
print(f"Temperature: {msg.output.temperature}")
42+
print(f"Conditions: {msg.output.conditions}")
43+
print(f"Humidity: {msg.output.humidity}%")
44+
print(f"Wind: {msg.output.wind_speed} mph")
45+
46+
47+
if __name__ == "__main__":
48+
asyncio.run(main())

src/vercel_ai_sdk/anthropic/__init__.py

Lines changed: 33 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66
from typing import Any, override
77

88
import anthropic
9+
import pydantic
910

1011
from .. import core
1112

@@ -119,6 +120,7 @@ async def stream_events(
119120
self,
120121
messages: list[core.messages.Message],
121122
tools: Sequence[core.tools.ToolLike] | None = None,
123+
output_type: type[pydantic.BaseModel] | None = None,
122124
) -> AsyncGenerator[core.llm.StreamEvent]:
123125
"""Yield raw stream events from Anthropic API."""
124126
system_prompt, anthropic_messages = _messages_to_anthropic(messages)
@@ -140,12 +142,30 @@ async def stream_events(
140142
"budget_tokens": self._budget_tokens,
141143
}
142144

145+
# Structured output: use beta API with output_format
146+
use_beta = False
147+
if output_type is not None:
148+
from anthropic.lib._parse._transform import transform_schema
149+
150+
kwargs["output_format"] = {
151+
"type": "json_schema",
152+
"schema": transform_schema(output_type),
153+
}
154+
kwargs["betas"] = ["structured-outputs-2025-11-13"]
155+
use_beta = True
156+
143157
# Track block types by index to know what End event to emit
144158
block_types: dict[int, str] = {} # index -> "text" | "thinking" | "tool_use"
145159
tool_ids: dict[int, str] = {} # index -> tool_call_id
146160
signature_buffer: dict[int, str] = {} # index -> accumulated signature
147161

148-
async with self._client.messages.stream(**kwargs) as stream:
162+
stream_cm: Any # BetaAsyncMessageStreamManager | AsyncMessageStreamManager
163+
if use_beta:
164+
stream_cm = self._client.beta.messages.stream(**kwargs)
165+
else:
166+
stream_cm = self._client.messages.stream(**kwargs)
167+
168+
async with stream_cm as stream:
149169
async for event in stream:
150170
if event.type == "content_block_start":
151171
block = event.content_block
@@ -208,8 +228,17 @@ async def stream(
208228
self,
209229
messages: list[core.messages.Message],
210230
tools: Sequence[core.tools.ToolLike] | None = None,
231+
output_type: type[pydantic.BaseModel] | None = None,
211232
) -> AsyncGenerator[core.messages.Message]:
212-
"""Stream Messages (uses StreamProcessor internally)."""
233+
"""Stream Messages (uses StreamHandler internally)."""
213234
handler = core.llm.StreamHandler()
214-
async for event in self.stream_events(messages, tools):
215-
yield handler.handle_event(event)
235+
msg: core.messages.Message | None = None
236+
async for event in self.stream_events(messages, tools, output_type=output_type):
237+
msg = handler.handle_event(event)
238+
yield msg
239+
240+
# After stream completes, validate and yield final message with output
241+
if output_type is not None and msg is not None and msg.text:
242+
msg = msg.model_copy()
243+
msg.output = output_type.model_validate_json(msg.text)
244+
yield msg

src/vercel_ai_sdk/core/llm.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,8 @@
44
import dataclasses
55
from collections.abc import AsyncGenerator, Sequence
66

7+
import pydantic
8+
79
from . import messages as messages_
810
from . import tools as tools_
911

@@ -216,6 +218,7 @@ async def stream(
216218
self,
217219
messages: list[messages_.Message],
218220
tools: Sequence[tools_.ToolLike] | None = None,
221+
output_type: type[pydantic.BaseModel] | None = None,
219222
) -> AsyncGenerator[messages_.Message]:
220223
raise NotImplementedError
221224
yield
@@ -224,10 +227,11 @@ async def buffer(
224227
self,
225228
messages: list[messages_.Message],
226229
tools: Sequence[tools_.ToolLike] | None = None,
230+
output_type: type[pydantic.BaseModel] | None = None,
227231
) -> messages_.Message:
228232
"""Drain the stream and return the final message."""
229233
final = None
230-
async for msg in self.stream(messages, tools):
234+
async for msg in self.stream(messages, tools, output_type=output_type):
231235
final = msg
232236
if final is None:
233237
raise ValueError("LLM produced no messages")

src/vercel_ai_sdk/core/messages.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -84,6 +84,7 @@ class Message(pydantic.BaseModel):
8484
parts: list[Part]
8585
id: str = pydantic.Field(default_factory=_gen_id)
8686
label: str | None = None
87+
output: Any = pydantic.Field(default=None, exclude=True)
8788

8889
@property
8990
def is_done(self) -> bool:

src/vercel_ai_sdk/core/streams.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,13 @@ def text(self) -> str:
3030
return self.last_message.text
3131
return ""
3232

33+
@property
34+
def output(self) -> Any:
35+
"""Parsed structured output from the last message, if available."""
36+
if self.last_message:
37+
return self.last_message.output
38+
return None
39+
3340

3441
Stream = Callable[[], AsyncGenerator[messages_.Message]]
3542
# maybe it should have a name and an id inferred from LLM outputs

src/vercel_ai_sdk/openai/__init__.py

Lines changed: 25 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
from typing import Any, override
66

77
import openai
8+
import pydantic
89

910
from .. import core
1011

@@ -138,6 +139,7 @@ async def stream_events(
138139
self,
139140
messages: list[core.messages.Message],
140141
tools: Sequence[core.tools.ToolLike] | None = None,
142+
output_type: type[pydantic.BaseModel] | None = None,
141143
) -> AsyncGenerator[core.llm.StreamEvent]:
142144
"""Yield raw stream events from OpenAI API."""
143145
openai_messages = _messages_to_openai(messages)
@@ -151,6 +153,18 @@ async def stream_events(
151153
if openai_tools:
152154
kwargs["tools"] = openai_tools
153155

156+
if output_type is not None:
157+
from openai.lib._pydantic import to_strict_json_schema
158+
159+
kwargs["response_format"] = {
160+
"type": "json_schema",
161+
"json_schema": {
162+
"name": output_type.__name__,
163+
"schema": to_strict_json_schema(output_type),
164+
"strict": True,
165+
},
166+
}
167+
154168
# Enable reasoning/thinking via Vercel AI Gateway's unified format
155169
# See: https://vercel.com/docs/ai-gateway/openai-compat/advanced
156170
if self._thinking:
@@ -249,8 +263,17 @@ async def stream(
249263
self,
250264
messages: list[core.messages.Message],
251265
tools: Sequence[core.tools.ToolLike] | None = None,
266+
output_type: type[pydantic.BaseModel] | None = None,
252267
) -> AsyncGenerator[core.messages.Message]:
253268
"""Stream Messages (uses StreamHandler internally)."""
254269
handler = core.llm.StreamHandler()
255-
async for event in self.stream_events(messages, tools):
256-
yield handler.handle_event(event)
270+
msg: core.messages.Message | None = None
271+
async for event in self.stream_events(messages, tools, output_type=output_type):
272+
msg = handler.handle_event(event)
273+
yield msg
274+
275+
# After stream completes, validate and yield final message with output
276+
if output_type is not None and msg is not None and msg.text:
277+
msg = msg.model_copy()
278+
msg.output = output_type.model_validate_json(msg.text)
279+
yield msg

tests/conftest.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,8 @@
22

33
from collections.abc import AsyncGenerator, Sequence
44

5+
import pydantic
6+
57
import vercel_ai_sdk as ai
68
from vercel_ai_sdk.core import messages
79

@@ -18,15 +20,23 @@ async def stream(
1820
self,
1921
messages: list[messages.Message],
2022
tools: Sequence[ai.ToolLike] | None = None,
23+
output_type: type[pydantic.BaseModel] | None = None,
2124
) -> AsyncGenerator[messages.Message]:
2225
if self._call_index >= len(self._responses):
2326
raise RuntimeError("MockLLM: no more responses configured")
2427
self.call_count += 1
2528
seq = self._responses[self._call_index]
2629
self._call_index += 1
30+
msg = None
2731
for msg in seq:
2832
yield msg
2933

34+
# Simulate structured output validation (matching real provider behavior)
35+
if output_type is not None and msg is not None and msg.text:
36+
msg = msg.model_copy()
37+
msg.output = output_type.model_validate_json(msg.text)
38+
yield msg
39+
3040

3141
def text_msg(
3242
text: str, *, id: str = "msg-1", state: str = "done", delta: str | None = None

0 commit comments

Comments
 (0)