saia-python 0.9.0__tar.gz → 0.10.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {saia_python-0.9.0/saia_python.egg-info → saia_python-0.10.0}/PKG-INFO +3 -2
- {saia_python-0.9.0 → saia_python-0.10.0}/README.md +1 -1
- {saia_python-0.9.0 → saia_python-0.10.0}/pyproject.toml +8 -1
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/__init__.py +12 -1
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/aio.py +34 -1
- saia_python-0.10.0/saia_python/chat.py +141 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/exceptions.py +26 -0
- saia_python-0.10.0/saia_python/structured.py +130 -0
- {saia_python-0.9.0 → saia_python-0.10.0/saia_python.egg-info}/PKG-INFO +3 -2
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python.egg-info/SOURCES.txt +4 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python.egg-info/requires.txt +1 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_async_streaming.py +24 -0
- saia_python-0.10.0/tests/test_live_responses_route.py +61 -0
- saia_python-0.10.0/tests/test_live_structured.py +50 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_streaming.py +24 -0
- saia_python-0.10.0/tests/test_structured.py +184 -0
- saia_python-0.9.0/saia_python/chat.py +0 -81
- {saia_python-0.9.0 → saia_python-0.10.0}/LICENSE +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/_async_http.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/_async_streaming.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/_http.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/_payloads.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/_streaming.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/_util.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/arcana.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/arcana_references.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/auth.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/client.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/documents.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/models.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/openai_compat.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/py.typed +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/rate_limits.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/responses.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/tokenizer.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python/voice.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python.egg-info/dependency_links.txt +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/saia_python.egg-info/top_level.txt +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/setup.cfg +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_arcana.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_arcana_references.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_async_arcana.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_async_chat.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_async_client.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_async_httpx_integration.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_async_transport.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_auth.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_chat.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_client.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_documents.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_exceptions.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_health_check.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_models.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_openai_compat.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_payloads.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_rate_limit_message.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_rate_limits.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_responses.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_setup_from_directory.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_tokenizer.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_transport_policy.py +0 -0
- {saia_python-0.9.0 → saia_python-0.10.0}/tests/test_voice.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: saia-python
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.10.0
|
|
4
4
|
Summary: Python wrapper for the GWDG SAIA platform REST API
|
|
5
5
|
Author: Friedrich Schwarz
|
|
6
6
|
License-Expression: AGPL-3.0-only
|
|
@@ -42,6 +42,7 @@ Requires-Dist: sentencepiece>=0.1.99; extra == "tokenizer"
|
|
|
42
42
|
Provides-Extra: test
|
|
43
43
|
Requires-Dist: pytest>=7.0; extra == "test"
|
|
44
44
|
Requires-Dist: pytest-cov>=4.0; extra == "test"
|
|
45
|
+
Requires-Dist: pydantic>=2; extra == "test"
|
|
45
46
|
Requires-Dist: saia-python[async,openai]; extra == "test"
|
|
46
47
|
Provides-Extra: docs
|
|
47
48
|
Requires-Dist: sphinx>=7.0; extra == "docs"
|
|
@@ -166,7 +167,7 @@ remain synchronous on `SAIAClient` — see
|
|
|
166
167
|
|
|
167
168
|
| Service | Description | GWDG Docs |
|
|
168
169
|
|---------|-------------|-----------|
|
|
169
|
-
| **Chat AI** | Chat completions with streaming
|
|
170
|
+
| **Chat AI** | Chat completions with streaming, tool calling, and structured output (Pydantic) | [Chat AI](https://docs.hpc.gwdg.de/services/ai-services/chat-ai/index.html) |
|
|
170
171
|
| **Voice AI** | Audio transcription and translation (Whisper) | [Voice AI](https://docs.hpc.gwdg.de/services/ai-services/voice-ai/index.html) |
|
|
171
172
|
| **ARCANA** | RAG — knowledge base management and retrieval-augmented chat | [ARCANA](https://docs.hpc.gwdg.de/services/ai-services/arcana/index.html) |
|
|
172
173
|
| **Documents** | PDF/document conversion via Docling | [SAIA API](https://docs.hpc.gwdg.de/services/ai-services/saia/index.html) |
|
|
@@ -106,7 +106,7 @@ remain synchronous on `SAIAClient` — see
|
|
|
106
106
|
|
|
107
107
|
| Service | Description | GWDG Docs |
|
|
108
108
|
|---------|-------------|-----------|
|
|
109
|
-
| **Chat AI** | Chat completions with streaming
|
|
109
|
+
| **Chat AI** | Chat completions with streaming, tool calling, and structured output (Pydantic) | [Chat AI](https://docs.hpc.gwdg.de/services/ai-services/chat-ai/index.html) |
|
|
110
110
|
| **Voice AI** | Audio transcription and translation (Whisper) | [Voice AI](https://docs.hpc.gwdg.de/services/ai-services/voice-ai/index.html) |
|
|
111
111
|
| **ARCANA** | RAG — knowledge base management and retrieval-augmented chat | [ARCANA](https://docs.hpc.gwdg.de/services/ai-services/arcana/index.html) |
|
|
112
112
|
| **Documents** | PDF/document conversion via Docling | [SAIA API](https://docs.hpc.gwdg.de/services/ai-services/saia/index.html) |
|
|
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
|
|
4
4
|
|
|
5
5
|
[project]
|
|
6
6
|
name = "saia-python"
|
|
7
|
-
version = "0.
|
|
7
|
+
version = "0.10.0"
|
|
8
8
|
description = "Python wrapper for the GWDG SAIA platform REST API"
|
|
9
9
|
readme = "README.md"
|
|
10
10
|
requires-python = ">=3.10"
|
|
@@ -81,6 +81,8 @@ tokenizer = [
|
|
|
81
81
|
test = [
|
|
82
82
|
"pytest>=7.0",
|
|
83
83
|
"pytest-cov>=4.0",
|
|
84
|
+
# the structured-output tests define Pydantic v2 models
|
|
85
|
+
"pydantic>=2",
|
|
84
86
|
"saia-python[openai,async]",
|
|
85
87
|
]
|
|
86
88
|
docs = [
|
|
@@ -132,3 +134,8 @@ disable_error_code = ["index", "operator", "union-attr", "call-arg"]
|
|
|
132
134
|
# `builtins.list`).
|
|
133
135
|
module = ["saia_python.models", "saia_python.arcana", "saia_python.aio"]
|
|
134
136
|
disable_error_code = ["valid-type"]
|
|
137
|
+
|
|
138
|
+
[dependency-groups]
|
|
139
|
+
dev = [
|
|
140
|
+
"ipykernel>=7.3.0",
|
|
141
|
+
]
|
|
@@ -44,10 +44,17 @@ from .auth import (
|
|
|
44
44
|
)
|
|
45
45
|
from .client import SAIAClient
|
|
46
46
|
from .documents import ConversionImage, ConversionResult
|
|
47
|
-
from .exceptions import
|
|
47
|
+
from .exceptions import (
|
|
48
|
+
APIError,
|
|
49
|
+
AuthenticationError,
|
|
50
|
+
RateLimitError,
|
|
51
|
+
SAIAError,
|
|
52
|
+
StructuredOutputError,
|
|
53
|
+
)
|
|
48
54
|
from .openai_compat import create_openai_client
|
|
49
55
|
from .rate_limits import RateLimitInfo, format_rate_limit_error, parse_rate_limits
|
|
50
56
|
from .responses import text_of
|
|
57
|
+
from .structured import parse_structured, response_format_for
|
|
51
58
|
from .tokenizer import (
|
|
52
59
|
DEFAULT_TOKENIZER_DIR,
|
|
53
60
|
GWDG_MODEL_REPOS,
|
|
@@ -98,6 +105,7 @@ __all__ = [
|
|
|
98
105
|
"AuthenticationError",
|
|
99
106
|
"RateLimitError",
|
|
100
107
|
"APIError",
|
|
108
|
+
"StructuredOutputError",
|
|
101
109
|
# Rate limits
|
|
102
110
|
"RateLimitInfo",
|
|
103
111
|
"parse_rate_limits",
|
|
@@ -110,6 +118,9 @@ __all__ = [
|
|
|
110
118
|
# Response helpers
|
|
111
119
|
"text_of",
|
|
112
120
|
"SSEStream",
|
|
121
|
+
# Structured output (Pydantic v2 models)
|
|
122
|
+
"response_format_for",
|
|
123
|
+
"parse_structured",
|
|
113
124
|
# ARCANA reference parsing
|
|
114
125
|
"ArcanaReference",
|
|
115
126
|
"ParsedReferences",
|
|
@@ -41,7 +41,7 @@ the two transports cannot drift.
|
|
|
41
41
|
|
|
42
42
|
from __future__ import annotations
|
|
43
43
|
|
|
44
|
-
from typing import TYPE_CHECKING, Any
|
|
44
|
+
from typing import TYPE_CHECKING, Any, cast
|
|
45
45
|
|
|
46
46
|
from ._async_http import aexecute, apost_chat_completion
|
|
47
47
|
from ._async_streaming import AsyncSSEStream
|
|
@@ -54,6 +54,7 @@ from ._payloads import (
|
|
|
54
54
|
from .auth import resolve_credentials
|
|
55
55
|
from .exceptions import raise_for_status
|
|
56
56
|
from .rate_limits import RateLimitInfo, parse_rate_limits
|
|
57
|
+
from .structured import ModelT, parse_structured, response_format_for
|
|
57
58
|
|
|
58
59
|
if TYPE_CHECKING:
|
|
59
60
|
import httpx
|
|
@@ -140,6 +141,38 @@ class AsyncChatService:
|
|
|
140
141
|
policy=resolve_retry(self._retry, retry),
|
|
141
142
|
)
|
|
142
143
|
|
|
144
|
+
async def completions_structured(
|
|
145
|
+
self,
|
|
146
|
+
model: str,
|
|
147
|
+
messages: list[dict],
|
|
148
|
+
response_model: type[ModelT],
|
|
149
|
+
*,
|
|
150
|
+
temperature: float | None = None,
|
|
151
|
+
top_p: float | None = None,
|
|
152
|
+
max_tokens: int | None = None,
|
|
153
|
+
retry: RetryPolicy | bool | None = None,
|
|
154
|
+
**kwargs: Any,
|
|
155
|
+
) -> ModelT:
|
|
156
|
+
"""Return the answer as a validated ``response_model`` instance.
|
|
157
|
+
|
|
158
|
+
See :meth:`ChatService.completions_structured
|
|
159
|
+
<saia_python.chat.ChatService.completions_structured>`; raises
|
|
160
|
+
:class:`~saia_python.StructuredOutputError` when there is no answer
|
|
161
|
+
that validates.
|
|
162
|
+
"""
|
|
163
|
+
response = await self.completions(
|
|
164
|
+
model,
|
|
165
|
+
messages,
|
|
166
|
+
temperature=temperature,
|
|
167
|
+
top_p=top_p,
|
|
168
|
+
max_tokens=max_tokens,
|
|
169
|
+
stream=False,
|
|
170
|
+
retry=retry,
|
|
171
|
+
response_format=response_format_for(response_model),
|
|
172
|
+
**kwargs,
|
|
173
|
+
)
|
|
174
|
+
return parse_structured(cast(dict, response), response_model)
|
|
175
|
+
|
|
143
176
|
def __repr__(self) -> str:
|
|
144
177
|
return f"AsyncChatService(base_url={self._base_url!r})"
|
|
145
178
|
|
|
@@ -0,0 +1,141 @@
|
|
|
1
|
+
"""Chat service — completions and streaming."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from typing import TYPE_CHECKING, cast
|
|
6
|
+
|
|
7
|
+
from ._http import RetryPolicy, coerce_retry, post_chat_completion, resolve_retry
|
|
8
|
+
from ._streaming import SSEStream
|
|
9
|
+
from .structured import ModelT, parse_structured, response_format_for
|
|
10
|
+
|
|
11
|
+
if TYPE_CHECKING:
|
|
12
|
+
import requests
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class ChatService:
|
|
16
|
+
"""Access the ``/chat/completions`` endpoint.
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
session: A :class:`requests.Session` with auth headers configured.
|
|
20
|
+
base_url: The SAIA API base URL.
|
|
21
|
+
"""
|
|
22
|
+
|
|
23
|
+
def __init__(
|
|
24
|
+
self,
|
|
25
|
+
session: requests.Session,
|
|
26
|
+
base_url: str,
|
|
27
|
+
*,
|
|
28
|
+
retry: RetryPolicy | bool | None = None,
|
|
29
|
+
):
|
|
30
|
+
self._session = session
|
|
31
|
+
self._base_url = base_url
|
|
32
|
+
self._retry = coerce_retry(retry)
|
|
33
|
+
|
|
34
|
+
def completions(
|
|
35
|
+
self,
|
|
36
|
+
model: str,
|
|
37
|
+
messages: list[dict],
|
|
38
|
+
*,
|
|
39
|
+
temperature: float | None = None,
|
|
40
|
+
top_p: float | None = None,
|
|
41
|
+
max_tokens: int | None = None,
|
|
42
|
+
stream: bool = False,
|
|
43
|
+
retry: RetryPolicy | bool | None = None,
|
|
44
|
+
**kwargs,
|
|
45
|
+
) -> dict | SSEStream:
|
|
46
|
+
"""Send a chat completion request.
|
|
47
|
+
|
|
48
|
+
Args:
|
|
49
|
+
model: Model identifier (e.g. ``"meta-llama-3.1-8b-instruct"``).
|
|
50
|
+
messages: List of message dicts with ``"role"`` and ``"content"`` keys.
|
|
51
|
+
temperature: Sampling temperature (0–2).
|
|
52
|
+
top_p: Nucleus sampling parameter (0–1).
|
|
53
|
+
max_tokens: Maximum tokens to generate.
|
|
54
|
+
stream: If ``True``, return a generator yielding chunks.
|
|
55
|
+
**kwargs: Additional parameters forwarded to the API.
|
|
56
|
+
|
|
57
|
+
Returns:
|
|
58
|
+
When ``stream=False``: the API response dict, with an extra
|
|
59
|
+
``"_rate_limits"`` key — a JSON-serializable dict of the current
|
|
60
|
+
rate-limit headers (see :class:`~saia_python.RateLimitInfo`).
|
|
61
|
+
When ``stream=True``: an ``SSEStream`` — iterate it for the
|
|
62
|
+
response chunks; its ``rate_limits`` attribute exposes the same
|
|
63
|
+
dict (available immediately, from the response headers).
|
|
64
|
+
"""
|
|
65
|
+
body = {"model": model, "messages": messages, **kwargs}
|
|
66
|
+
if temperature is not None:
|
|
67
|
+
body["temperature"] = temperature
|
|
68
|
+
if top_p is not None:
|
|
69
|
+
body["top_p"] = top_p
|
|
70
|
+
if max_tokens is not None:
|
|
71
|
+
body["max_tokens"] = max_tokens
|
|
72
|
+
|
|
73
|
+
return post_chat_completion(
|
|
74
|
+
self._session,
|
|
75
|
+
f"{self._base_url}/chat/completions",
|
|
76
|
+
body,
|
|
77
|
+
stream=stream,
|
|
78
|
+
policy=resolve_retry(self._retry, retry),
|
|
79
|
+
)
|
|
80
|
+
|
|
81
|
+
def completions_structured(
|
|
82
|
+
self,
|
|
83
|
+
model: str,
|
|
84
|
+
messages: list[dict],
|
|
85
|
+
response_model: type[ModelT],
|
|
86
|
+
*,
|
|
87
|
+
temperature: float | None = None,
|
|
88
|
+
top_p: float | None = None,
|
|
89
|
+
max_tokens: int | None = None,
|
|
90
|
+
retry: RetryPolicy | bool | None = None,
|
|
91
|
+
**kwargs,
|
|
92
|
+
) -> ModelT:
|
|
93
|
+
"""Send a chat completion and return the answer as a validated model.
|
|
94
|
+
|
|
95
|
+
Sends ``response_model``'s JSON Schema as the ``response_format`` (see
|
|
96
|
+
:func:`~saia_python.response_format_for`), which SAIA enforces on the
|
|
97
|
+
server, then validates the answer into an instance (see
|
|
98
|
+
:func:`~saia_python.parse_structured`). Non-streaming only. To keep the
|
|
99
|
+
raw response too (``usage``, ``_rate_limits``), call those two helpers
|
|
100
|
+
around :meth:`completions` yourself.
|
|
101
|
+
|
|
102
|
+
Args:
|
|
103
|
+
model: Model identifier (e.g. ``"meta-llama-3.1-8b-instruct"``).
|
|
104
|
+
messages: List of message dicts with ``"role"`` and ``"content"`` keys.
|
|
105
|
+
SAIA uses the schema only to constrain the output and does not
|
|
106
|
+
add it to the prompt, so say what each field should hold.
|
|
107
|
+
response_model: The Pydantic v2 model class the answer must match.
|
|
108
|
+
temperature: Sampling temperature (0–2).
|
|
109
|
+
top_p: Nucleus sampling parameter (0–1).
|
|
110
|
+
max_tokens: Maximum tokens to generate. Reasoning models think
|
|
111
|
+
before they answer, so leave room, or turn thinking off where
|
|
112
|
+
the chat template allows it (e.g. Qwen:
|
|
113
|
+
``chat_template_kwargs={"enable_thinking": False}``).
|
|
114
|
+
retry: Overrides the service's rate-limit retry policy for this
|
|
115
|
+
call, as in :meth:`completions`.
|
|
116
|
+
**kwargs: Additional parameters forwarded to the API.
|
|
117
|
+
|
|
118
|
+
Returns:
|
|
119
|
+
An instance of ``response_model``.
|
|
120
|
+
|
|
121
|
+
Raises:
|
|
122
|
+
StructuredOutputError: The response holds no answer that validates,
|
|
123
|
+
e.g. because the token budget ran out (``finish_reason`` is
|
|
124
|
+
``"length"``). The error keeps the full response, ``usage``
|
|
125
|
+
included.
|
|
126
|
+
"""
|
|
127
|
+
response = self.completions(
|
|
128
|
+
model,
|
|
129
|
+
messages,
|
|
130
|
+
temperature=temperature,
|
|
131
|
+
top_p=top_p,
|
|
132
|
+
max_tokens=max_tokens,
|
|
133
|
+
stream=False,
|
|
134
|
+
retry=retry,
|
|
135
|
+
response_format=response_format_for(response_model),
|
|
136
|
+
**kwargs,
|
|
137
|
+
)
|
|
138
|
+
return parse_structured(cast(dict, response), response_model)
|
|
139
|
+
|
|
140
|
+
def __repr__(self):
|
|
141
|
+
return f"ChatService(base_url={self._base_url!r})"
|
|
@@ -76,6 +76,32 @@ class APIError(SAIAError):
|
|
|
76
76
|
super().__init__(message, status_code=status_code, response_body=response_body)
|
|
77
77
|
|
|
78
78
|
|
|
79
|
+
class StructuredOutputError(SAIAError):
|
|
80
|
+
"""Raised when a chat response holds no answer that validates against the
|
|
81
|
+
requested structured-output model.
|
|
82
|
+
|
|
83
|
+
The HTTP call itself succeeded, so ``status_code`` is ``None``; the message
|
|
84
|
+
says why the answer is unusable. The response stays attached, so a caller
|
|
85
|
+
can still read ``usage`` (the tokens were spent) or inspect the raw content.
|
|
86
|
+
|
|
87
|
+
Attributes:
|
|
88
|
+
response: The chat response dict the answer was read from.
|
|
89
|
+
finish_reason: The first choice's ``finish_reason`` — ``"length"`` means
|
|
90
|
+
the token budget ran out — or ``None`` when there was no choice.
|
|
91
|
+
"""
|
|
92
|
+
|
|
93
|
+
def __init__(
|
|
94
|
+
self,
|
|
95
|
+
message: object,
|
|
96
|
+
*,
|
|
97
|
+
response: dict,
|
|
98
|
+
finish_reason: str | None = None,
|
|
99
|
+
):
|
|
100
|
+
super().__init__(message)
|
|
101
|
+
self.response = response
|
|
102
|
+
self.finish_reason = finish_reason
|
|
103
|
+
|
|
104
|
+
|
|
79
105
|
def _extract_detail(resp: Any) -> str:
|
|
80
106
|
"""Try to extract a human-readable message from a JSON error body.
|
|
81
107
|
|
|
@@ -0,0 +1,130 @@
|
|
|
1
|
+
"""Structured output — a Pydantic model in, a validated instance out.
|
|
2
|
+
|
|
3
|
+
SAIA enforces a ``json_schema`` ``response_format`` on the server: its inference
|
|
4
|
+
backend (vLLM) compiles the schema into a grammar and blocks every token that
|
|
5
|
+
would break it, so the answer matches the schema by construction and needs no
|
|
6
|
+
validate-and-retry loop. This module covers the two client-side ends:
|
|
7
|
+
|
|
8
|
+
- :func:`response_format_for` turns a Pydantic model into the
|
|
9
|
+
``response_format`` request field.
|
|
10
|
+
- :func:`parse_structured` validates a chat response back into the model, and
|
|
11
|
+
raises :class:`~saia_python.StructuredOutputError` with the reason when there
|
|
12
|
+
is nothing usable — most often a reasoning model that spent its
|
|
13
|
+
``max_tokens`` thinking.
|
|
14
|
+
|
|
15
|
+
:meth:`ChatService.completions_structured
|
|
16
|
+
<saia_python.chat.ChatService.completions_structured>` does both in one call.
|
|
17
|
+
Use the two helpers directly when you also need the raw response, e.g. its
|
|
18
|
+
``usage``::
|
|
19
|
+
|
|
20
|
+
resp = client.chat.completions(
|
|
21
|
+
model, messages, response_format=response_format_for(Invoice)
|
|
22
|
+
)
|
|
23
|
+
invoice = parse_structured(resp, Invoice)
|
|
24
|
+
tokens = resp["usage"]["total_tokens"]
|
|
25
|
+
|
|
26
|
+
Needs Pydantic v2, which you already have if you define a model; the package
|
|
27
|
+
itself imports it only inside :func:`parse_structured`.
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
from __future__ import annotations
|
|
31
|
+
|
|
32
|
+
import re
|
|
33
|
+
from typing import TYPE_CHECKING, TypeVar
|
|
34
|
+
|
|
35
|
+
from .exceptions import StructuredOutputError
|
|
36
|
+
|
|
37
|
+
if TYPE_CHECKING:
|
|
38
|
+
from pydantic import BaseModel
|
|
39
|
+
|
|
40
|
+
ModelT = TypeVar("ModelT", bound="BaseModel")
|
|
41
|
+
|
|
42
|
+
_TOKEN_HINT = (
|
|
43
|
+
"Reasoning models spend tokens thinking before they answer: raise max_tokens, "
|
|
44
|
+
"or turn thinking off where the chat template allows it (e.g. Qwen: "
|
|
45
|
+
"chat_template_kwargs={'enable_thinking': False})."
|
|
46
|
+
)
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def response_format_for(response_model: type[BaseModel]) -> dict:
|
|
50
|
+
"""Build the ``response_format`` request field for a Pydantic model.
|
|
51
|
+
|
|
52
|
+
``strict`` is left unset: SAIA's backends enforce the schema either way,
|
|
53
|
+
whereas OpenAI's strict mode would reject a plain Pydantic schema (it
|
|
54
|
+
requires ``additionalProperties: false`` on every object).
|
|
55
|
+
|
|
56
|
+
Args:
|
|
57
|
+
response_model: A Pydantic v2 model class.
|
|
58
|
+
|
|
59
|
+
Returns:
|
|
60
|
+
``{"type": "json_schema", "json_schema": {"name": ..., "schema": ...}}``
|
|
61
|
+
with the model's JSON Schema, named after the class (characters outside
|
|
62
|
+
``a-z A-Z 0-9 _ -``, such as a generic's brackets, become ``_``).
|
|
63
|
+
"""
|
|
64
|
+
name = re.sub(r"[^a-zA-Z0-9_-]", "_", response_model.__name__)[:64]
|
|
65
|
+
return {
|
|
66
|
+
"type": "json_schema",
|
|
67
|
+
"json_schema": {"name": name, "schema": response_model.model_json_schema()},
|
|
68
|
+
}
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def parse_structured(response: dict, response_model: type[ModelT]) -> ModelT:
|
|
72
|
+
"""Validate a chat response's answer into ``response_model``.
|
|
73
|
+
|
|
74
|
+
Reads the first choice's ``content`` — the JSON string the model wrote — and
|
|
75
|
+
validates it with ``response_model.model_validate_json``. Works on any
|
|
76
|
+
OpenAI-style ChatCompletion dict, such as the one
|
|
77
|
+
:meth:`~saia_python.chat.ChatService.completions` returns.
|
|
78
|
+
|
|
79
|
+
Args:
|
|
80
|
+
response: A chat completion response dict.
|
|
81
|
+
response_model: The Pydantic v2 model class to validate into.
|
|
82
|
+
|
|
83
|
+
Returns:
|
|
84
|
+
The validated ``response_model`` instance.
|
|
85
|
+
|
|
86
|
+
Raises:
|
|
87
|
+
StructuredOutputError: When there is no usable answer: no choices; no
|
|
88
|
+
content (a reasoning model that ran out of tokens while thinking
|
|
89
|
+
returns ``content: None`` with ``finish_reason: "length"``); an
|
|
90
|
+
answer cut off at ``max_tokens``; or one that fails validation, with
|
|
91
|
+
the :class:`pydantic.ValidationError` chained as the cause.
|
|
92
|
+
"""
|
|
93
|
+
from pydantic import ValidationError
|
|
94
|
+
|
|
95
|
+
name = response_model.__name__
|
|
96
|
+
choices = response.get("choices") or []
|
|
97
|
+
if not choices:
|
|
98
|
+
raise StructuredOutputError(
|
|
99
|
+
f"{name}: the response has no choices", response=response
|
|
100
|
+
)
|
|
101
|
+
finish_reason = choices[0].get("finish_reason")
|
|
102
|
+
message = choices[0].get("message") or {}
|
|
103
|
+
content = message.get("content")
|
|
104
|
+
if not content:
|
|
105
|
+
if finish_reason == "length":
|
|
106
|
+
reason = f"the model ran out of tokens before answering. {_TOKEN_HINT}"
|
|
107
|
+
elif message.get("refusal"):
|
|
108
|
+
reason = f"the model refused: {message['refusal']}"
|
|
109
|
+
else:
|
|
110
|
+
reason = f"the response has no content (finish_reason={finish_reason!r})"
|
|
111
|
+
raise StructuredOutputError(
|
|
112
|
+
f"{name}: {reason}", response=response, finish_reason=finish_reason
|
|
113
|
+
)
|
|
114
|
+
try:
|
|
115
|
+
return response_model.model_validate_json(content)
|
|
116
|
+
except ValidationError as exc:
|
|
117
|
+
# Validate before blaming finish_reason: an answer can be complete even
|
|
118
|
+
# when generation hit the limit afterwards (e.g. trailing whitespace).
|
|
119
|
+
if finish_reason == "length":
|
|
120
|
+
reason = f"the answer was cut off at max_tokens. {_TOKEN_HINT}"
|
|
121
|
+
else:
|
|
122
|
+
error = exc.errors()[0]
|
|
123
|
+
where = ".".join(str(part) for part in error["loc"])
|
|
124
|
+
detail = f"{where}: {error['msg']}" if where else error["msg"]
|
|
125
|
+
if exc.error_count() > 1:
|
|
126
|
+
detail += f"; {exc.error_count() - 1} more"
|
|
127
|
+
reason = f"the answer failed validation ({detail})"
|
|
128
|
+
raise StructuredOutputError(
|
|
129
|
+
f"{name}: {reason}", response=response, finish_reason=finish_reason
|
|
130
|
+
) from exc
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: saia-python
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.10.0
|
|
4
4
|
Summary: Python wrapper for the GWDG SAIA platform REST API
|
|
5
5
|
Author: Friedrich Schwarz
|
|
6
6
|
License-Expression: AGPL-3.0-only
|
|
@@ -42,6 +42,7 @@ Requires-Dist: sentencepiece>=0.1.99; extra == "tokenizer"
|
|
|
42
42
|
Provides-Extra: test
|
|
43
43
|
Requires-Dist: pytest>=7.0; extra == "test"
|
|
44
44
|
Requires-Dist: pytest-cov>=4.0; extra == "test"
|
|
45
|
+
Requires-Dist: pydantic>=2; extra == "test"
|
|
45
46
|
Requires-Dist: saia-python[async,openai]; extra == "test"
|
|
46
47
|
Provides-Extra: docs
|
|
47
48
|
Requires-Dist: sphinx>=7.0; extra == "docs"
|
|
@@ -166,7 +167,7 @@ remain synchronous on `SAIAClient` — see
|
|
|
166
167
|
|
|
167
168
|
| Service | Description | GWDG Docs |
|
|
168
169
|
|---------|-------------|-----------|
|
|
169
|
-
| **Chat AI** | Chat completions with streaming
|
|
170
|
+
| **Chat AI** | Chat completions with streaming, tool calling, and structured output (Pydantic) | [Chat AI](https://docs.hpc.gwdg.de/services/ai-services/chat-ai/index.html) |
|
|
170
171
|
| **Voice AI** | Audio transcription and translation (Whisper) | [Voice AI](https://docs.hpc.gwdg.de/services/ai-services/voice-ai/index.html) |
|
|
171
172
|
| **ARCANA** | RAG — knowledge base management and retrieval-augmented chat | [ARCANA](https://docs.hpc.gwdg.de/services/ai-services/arcana/index.html) |
|
|
172
173
|
| **Documents** | PDF/document conversion via Docling | [SAIA API](https://docs.hpc.gwdg.de/services/ai-services/saia/index.html) |
|
|
@@ -21,6 +21,7 @@ saia_python/openai_compat.py
|
|
|
21
21
|
saia_python/py.typed
|
|
22
22
|
saia_python/rate_limits.py
|
|
23
23
|
saia_python/responses.py
|
|
24
|
+
saia_python/structured.py
|
|
24
25
|
saia_python/tokenizer.py
|
|
25
26
|
saia_python/voice.py
|
|
26
27
|
saia_python.egg-info/PKG-INFO
|
|
@@ -42,6 +43,8 @@ tests/test_client.py
|
|
|
42
43
|
tests/test_documents.py
|
|
43
44
|
tests/test_exceptions.py
|
|
44
45
|
tests/test_health_check.py
|
|
46
|
+
tests/test_live_responses_route.py
|
|
47
|
+
tests/test_live_structured.py
|
|
45
48
|
tests/test_models.py
|
|
46
49
|
tests/test_openai_compat.py
|
|
47
50
|
tests/test_payloads.py
|
|
@@ -50,6 +53,7 @@ tests/test_rate_limits.py
|
|
|
50
53
|
tests/test_responses.py
|
|
51
54
|
tests/test_setup_from_directory.py
|
|
52
55
|
tests/test_streaming.py
|
|
56
|
+
tests/test_structured.py
|
|
53
57
|
tests/test_tokenizer.py
|
|
54
58
|
tests/test_transport_policy.py
|
|
55
59
|
tests/test_voice.py
|
|
@@ -8,6 +8,7 @@ loop.
|
|
|
8
8
|
from __future__ import annotations
|
|
9
9
|
|
|
10
10
|
import asyncio
|
|
11
|
+
import json
|
|
11
12
|
|
|
12
13
|
import pytest
|
|
13
14
|
|
|
@@ -140,3 +141,26 @@ def test_aclose_is_idempotent():
|
|
|
140
141
|
|
|
141
142
|
asyncio.run(_run())
|
|
142
143
|
assert client.closed == 1
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
def test_keeps_trailing_usage_chunk_with_empty_choices():
|
|
147
|
+
# Same contract as the sync twin: SAIA's final chunk has no choices but
|
|
148
|
+
# carries the token usage, and the stream must hand it over unchanged.
|
|
149
|
+
usage_chunk = {
|
|
150
|
+
"choices": [],
|
|
151
|
+
"usage": {"prompt_tokens": 104, "total_tokens": 124, "completion_tokens": 20},
|
|
152
|
+
}
|
|
153
|
+
lines = [
|
|
154
|
+
'data: {"choices": [{"delta": {"content": "Hi"}}]}',
|
|
155
|
+
f"data: {json.dumps(usage_chunk)}",
|
|
156
|
+
"data: [DONE]",
|
|
157
|
+
]
|
|
158
|
+
client = FakeAsyncClient(
|
|
159
|
+
stream_responses=[FakeAsyncResponse(200, headers=rl_headers(), lines=lines)]
|
|
160
|
+
)
|
|
161
|
+
|
|
162
|
+
async def _run():
|
|
163
|
+
stream = await _open(client)
|
|
164
|
+
return [chunk async for chunk in stream]
|
|
165
|
+
|
|
166
|
+
assert asyncio.run(_run())[-1] == usage_chunk
|
|
@@ -0,0 +1,61 @@
|
|
|
1
|
+
"""Live check that SAIA's undocumented ``/v1/responses`` route still answers.
|
|
2
|
+
|
|
3
|
+
GWDG's docs say SAIA does not provide OpenAI's Responses API, yet the route
|
|
4
|
+
currently returns genuine Responses API objects. This test tracks that against
|
|
5
|
+
the real service, so it needs an API key and the network. It is gated behind the
|
|
6
|
+
``SAIA_RESPONSES_LIVE`` environment variable and never runs in CI or the offline
|
|
7
|
+
gate by default::
|
|
8
|
+
|
|
9
|
+
SAIA_RESPONSES_LIVE=1 pytest tests/test_live_responses_route.py
|
|
10
|
+
|
|
11
|
+
``SAIA_RESPONSES_LIVE_MODEL`` overrides the probed model. Keep it a
|
|
12
|
+
non-reasoning model: a reasoning model can spend the small token budget thinking
|
|
13
|
+
and come back ``incomplete``.
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
from __future__ import annotations
|
|
17
|
+
|
|
18
|
+
import os
|
|
19
|
+
|
|
20
|
+
import pytest
|
|
21
|
+
import requests
|
|
22
|
+
|
|
23
|
+
from saia_python import load_api_key, resolve_base_url
|
|
24
|
+
|
|
25
|
+
MODEL = os.environ.get("SAIA_RESPONSES_LIVE_MODEL", "meta-llama-3.1-8b-instruct")
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
@pytest.mark.skipif(
|
|
29
|
+
not os.environ.get("SAIA_RESPONSES_LIVE"),
|
|
30
|
+
reason="set SAIA_RESPONSES_LIVE=1 to run the network-backed /v1/responses check",
|
|
31
|
+
)
|
|
32
|
+
def test_live_responses_route_still_answers():
|
|
33
|
+
resp = requests.post(
|
|
34
|
+
f"{resolve_base_url()}/responses",
|
|
35
|
+
headers={"Authorization": f"Bearer {load_api_key()}"},
|
|
36
|
+
json={
|
|
37
|
+
"model": MODEL,
|
|
38
|
+
"instructions": "Answer in one word.",
|
|
39
|
+
"input": "What is the capital of France?",
|
|
40
|
+
"max_output_tokens": 16,
|
|
41
|
+
},
|
|
42
|
+
timeout=(10, 60),
|
|
43
|
+
)
|
|
44
|
+
if resp.status_code == 429:
|
|
45
|
+
pytest.skip("rate-limited (HTTP 429); cannot tell whether the route works")
|
|
46
|
+
assert resp.status_code == 200, f"HTTP {resp.status_code}: {resp.text[:300]}"
|
|
47
|
+
|
|
48
|
+
body = resp.json()
|
|
49
|
+
# Only the Responses API answers in this shape: Chat Completions would say
|
|
50
|
+
# "chat.completion", answer under "choices", and count prompt/completion tokens.
|
|
51
|
+
assert body["object"] == "response"
|
|
52
|
+
assert body["status"] == "completed"
|
|
53
|
+
texts = [
|
|
54
|
+
part["text"]
|
|
55
|
+
for item in body["output"]
|
|
56
|
+
if item.get("type") == "message"
|
|
57
|
+
for part in item.get("content", [])
|
|
58
|
+
if part.get("type") == "output_text"
|
|
59
|
+
]
|
|
60
|
+
assert any(text.strip() for text in texts), body["output"]
|
|
61
|
+
assert {"input_tokens", "output_tokens"} <= body["usage"].keys()
|
|
@@ -0,0 +1,50 @@
|
|
|
1
|
+
"""Live check that SAIA still enforces a ``json_schema`` ``response_format``.
|
|
2
|
+
|
|
3
|
+
Calls :meth:`ChatService.completions_structured` against the real service, so it
|
|
4
|
+
needs an API key and the network. It is gated behind the
|
|
5
|
+
``SAIA_STRUCTURED_LIVE`` environment variable and never runs in CI or the
|
|
6
|
+
offline gate by default::
|
|
7
|
+
|
|
8
|
+
SAIA_STRUCTURED_LIVE=1 pytest tests/test_live_structured.py
|
|
9
|
+
|
|
10
|
+
The prompt never mentions JSON, so a schema-shaped answer shows that the server
|
|
11
|
+
enforces the schema, not that the model followed instructions.
|
|
12
|
+
``SAIA_STRUCTURED_LIVE_MODEL`` overrides the model. No ``max_tokens`` is set, so
|
|
13
|
+
a reasoning model has room to think before it answers.
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
from __future__ import annotations
|
|
17
|
+
|
|
18
|
+
import os
|
|
19
|
+
from typing import Literal
|
|
20
|
+
|
|
21
|
+
import pytest
|
|
22
|
+
from pydantic import BaseModel
|
|
23
|
+
|
|
24
|
+
from saia_python import RateLimitError, SAIAClient
|
|
25
|
+
|
|
26
|
+
MODEL = os.environ.get("SAIA_STRUCTURED_LIVE_MODEL", "meta-llama-3.1-8b-instruct")
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class Sentiment(BaseModel):
|
|
30
|
+
label: Literal["positive", "negative", "neutral"]
|
|
31
|
+
confidence: float
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
@pytest.mark.skipif(
|
|
35
|
+
not os.environ.get("SAIA_STRUCTURED_LIVE"),
|
|
36
|
+
reason="set SAIA_STRUCTURED_LIVE=1 to run the network-backed structured-output check",
|
|
37
|
+
)
|
|
38
|
+
def test_live_completions_structured_returns_a_validated_model():
|
|
39
|
+
try:
|
|
40
|
+
result = SAIAClient().chat.completions_structured(
|
|
41
|
+
MODEL,
|
|
42
|
+
[{"role": "user", "content": "Sentiment of: 'The update broke my build.'"}],
|
|
43
|
+
Sentiment,
|
|
44
|
+
retry=False,
|
|
45
|
+
)
|
|
46
|
+
except RateLimitError:
|
|
47
|
+
pytest.skip("rate-limited (HTTP 429); cannot tell whether the schema holds")
|
|
48
|
+
# completions_structured raises StructuredOutputError on anything unusable,
|
|
49
|
+
# so reaching this line means the answer validated.
|
|
50
|
+
assert isinstance(result, Sentiment)
|
|
@@ -1,5 +1,6 @@
|
|
|
1
1
|
"""Tests for saia_python._streaming — SSE line parsing."""
|
|
2
2
|
|
|
3
|
+
import json
|
|
3
4
|
from unittest.mock import MagicMock
|
|
4
5
|
|
|
5
6
|
import pytest
|
|
@@ -98,3 +99,26 @@ class TestSSEStream:
|
|
|
98
99
|
stream = SSEStream(resp)
|
|
99
100
|
stream.close()
|
|
100
101
|
resp.close.assert_called()
|
|
102
|
+
|
|
103
|
+
def test_keeps_trailing_usage_chunk_with_empty_choices(self):
|
|
104
|
+
# SAIA ends every chat stream with a chunk that has no choices but carries
|
|
105
|
+
# the token usage, even without stream_options.include_usage. Callers bill
|
|
106
|
+
# from it, so the stream must hand it over unchanged.
|
|
107
|
+
usage_chunk = {
|
|
108
|
+
"choices": [],
|
|
109
|
+
"usage": {
|
|
110
|
+
"prompt_tokens": 104,
|
|
111
|
+
"total_tokens": 124,
|
|
112
|
+
"completion_tokens": 20,
|
|
113
|
+
},
|
|
114
|
+
}
|
|
115
|
+
resp = _make_response(
|
|
116
|
+
[
|
|
117
|
+
'data: {"choices": [{"delta": {"content": "Hi"}}]}',
|
|
118
|
+
f"data: {json.dumps(usage_chunk)}",
|
|
119
|
+
"data: [DONE]",
|
|
120
|
+
]
|
|
121
|
+
)
|
|
122
|
+
resp.headers = {}
|
|
123
|
+
chunks = list(SSEStream(resp))
|
|
124
|
+
assert chunks[-1] == usage_chunk
|
|
@@ -0,0 +1,184 @@
|
|
|
1
|
+
"""Tests for structured output — ``response_format_for``, ``parse_structured``,
|
|
2
|
+
and ``completions_structured`` on the sync and async chat services.
|
|
3
|
+
|
|
4
|
+
The error cases mirror what SAIA returned live (2026-10-08): a reasoning model
|
|
5
|
+
that runs out of tokens while thinking answers HTTP 200 with ``content: None``
|
|
6
|
+
and ``finish_reason: "length"``.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import asyncio
|
|
12
|
+
from typing import Generic, Literal, TypeVar
|
|
13
|
+
from unittest.mock import MagicMock
|
|
14
|
+
|
|
15
|
+
import pytest
|
|
16
|
+
from pydantic import BaseModel, ValidationError
|
|
17
|
+
|
|
18
|
+
from saia_python import (
|
|
19
|
+
SAIAError,
|
|
20
|
+
StructuredOutputError,
|
|
21
|
+
parse_structured,
|
|
22
|
+
response_format_for,
|
|
23
|
+
)
|
|
24
|
+
from saia_python._http import RetryPolicy
|
|
25
|
+
from saia_python.aio import AsyncChatService
|
|
26
|
+
from saia_python.chat import ChatService
|
|
27
|
+
|
|
28
|
+
from ._async_fakes import FakeAsyncClient, FakeAsyncResponse, rl_headers
|
|
29
|
+
|
|
30
|
+
T = TypeVar("T")
|
|
31
|
+
|
|
32
|
+
GOOD = '{"label": "negative", "confidence": 0.95}'
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
class Sentiment(BaseModel):
|
|
36
|
+
label: Literal["positive", "negative", "neutral"]
|
|
37
|
+
confidence: float
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
class Wrapper(BaseModel, Generic[T]):
|
|
41
|
+
value: T
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def _response(content, finish_reason="stop", **message):
|
|
45
|
+
return {
|
|
46
|
+
"choices": [
|
|
47
|
+
{
|
|
48
|
+
"message": {"role": "assistant", "content": content, **message},
|
|
49
|
+
"finish_reason": finish_reason,
|
|
50
|
+
}
|
|
51
|
+
],
|
|
52
|
+
"usage": {"prompt_tokens": 21, "completion_tokens": 150, "total_tokens": 171},
|
|
53
|
+
}
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def _sync_service(json_body) -> ChatService:
|
|
57
|
+
svc = ChatService.__new__(ChatService)
|
|
58
|
+
svc._session = MagicMock()
|
|
59
|
+
svc._base_url = "https://example.com/v1"
|
|
60
|
+
svc._retry = RetryPolicy()
|
|
61
|
+
resp = MagicMock()
|
|
62
|
+
resp.ok = True
|
|
63
|
+
resp.status_code = 200
|
|
64
|
+
resp.headers = rl_headers()
|
|
65
|
+
resp.json.return_value = json_body
|
|
66
|
+
svc._session.post.return_value = resp
|
|
67
|
+
return svc
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def test_response_format_for_sends_the_model_schema():
|
|
71
|
+
assert response_format_for(Sentiment) == {
|
|
72
|
+
"type": "json_schema",
|
|
73
|
+
"json_schema": {"name": "Sentiment", "schema": Sentiment.model_json_schema()},
|
|
74
|
+
}
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
def test_response_format_for_sanitizes_generic_model_names():
|
|
78
|
+
# OpenAI-style APIs allow only a-z, A-Z, 0-9, _ and - in the name.
|
|
79
|
+
assert response_format_for(Wrapper[int])["json_schema"]["name"] == "Wrapper_int_"
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
def test_parse_structured_returns_validated_instance():
|
|
83
|
+
result = parse_structured(_response(GOOD), Sentiment)
|
|
84
|
+
assert result == Sentiment(label="negative", confidence=0.95)
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
def test_reasoning_model_out_of_tokens_raises_with_hint():
|
|
88
|
+
resp = _response(None, finish_reason="length", reasoning="Okay, the user...")
|
|
89
|
+
with pytest.raises(StructuredOutputError) as info:
|
|
90
|
+
parse_structured(resp, Sentiment)
|
|
91
|
+
err = info.value
|
|
92
|
+
assert isinstance(err, SAIAError)
|
|
93
|
+
assert err.finish_reason == "length"
|
|
94
|
+
assert "max_tokens" in str(err)
|
|
95
|
+
assert "enable_thinking" in str(err)
|
|
96
|
+
# The tokens were spent, so a caller billing from usage can still read it.
|
|
97
|
+
assert err.response["usage"]["total_tokens"] == 171
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
def test_answer_cut_off_at_max_tokens_raises_truncation_error():
|
|
101
|
+
resp = _response('{"label": "nega', finish_reason="length")
|
|
102
|
+
with pytest.raises(StructuredOutputError, match="cut off at max_tokens") as info:
|
|
103
|
+
parse_structured(resp, Sentiment)
|
|
104
|
+
assert isinstance(info.value.__cause__, ValidationError)
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def test_complete_answer_is_kept_even_when_generation_hit_the_limit():
|
|
108
|
+
resp = _response(GOOD + "\n\n\n", finish_reason="length")
|
|
109
|
+
assert parse_structured(resp, Sentiment).label == "negative"
|
|
110
|
+
|
|
111
|
+
|
|
112
|
+
def test_mismatching_answer_names_the_failing_field():
|
|
113
|
+
# Bad label and missing confidence: the first error is named, the rest counted.
|
|
114
|
+
resp = _response('{"label": "great"}')
|
|
115
|
+
with pytest.raises(StructuredOutputError) as info:
|
|
116
|
+
parse_structured(resp, Sentiment)
|
|
117
|
+
message = str(info.value)
|
|
118
|
+
assert "failed validation (label: " in message
|
|
119
|
+
assert message.endswith("; 1 more)")
|
|
120
|
+
assert info.value.finish_reason == "stop"
|
|
121
|
+
assert isinstance(info.value.__cause__, ValidationError)
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
def test_non_json_answer_reports_the_parse_error():
|
|
125
|
+
# What a backend that ignores the schema might send: JSON in a code fence.
|
|
126
|
+
resp = _response(f"```json\n{GOOD}\n```")
|
|
127
|
+
with pytest.raises(
|
|
128
|
+
StructuredOutputError, match=r"failed validation \(Invalid JSON"
|
|
129
|
+
):
|
|
130
|
+
parse_structured(resp, Sentiment)
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
def test_missing_content_reports_the_finish_reason():
|
|
134
|
+
# E.g. the model called a tool instead of answering.
|
|
135
|
+
resp = _response(None, finish_reason="tool_calls")
|
|
136
|
+
with pytest.raises(StructuredOutputError, match="finish_reason='tool_calls'"):
|
|
137
|
+
parse_structured(resp, Sentiment)
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
def test_refusal_is_reported():
|
|
141
|
+
resp = _response(None, refusal="I can't help with that.")
|
|
142
|
+
with pytest.raises(StructuredOutputError, match="refused: I can't help with that"):
|
|
143
|
+
parse_structured(resp, Sentiment)
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
def test_no_choices_raises():
|
|
147
|
+
with pytest.raises(StructuredOutputError, match="no choices") as info:
|
|
148
|
+
parse_structured({"choices": []}, Sentiment)
|
|
149
|
+
assert info.value.finish_reason is None
|
|
150
|
+
|
|
151
|
+
|
|
152
|
+
def test_completions_structured_sends_schema_and_returns_instance():
|
|
153
|
+
svc = _sync_service(_response(GOOD))
|
|
154
|
+
result = svc.completions_structured(
|
|
155
|
+
"m",
|
|
156
|
+
[{"role": "user", "content": "hi"}],
|
|
157
|
+
Sentiment,
|
|
158
|
+
max_tokens=2000,
|
|
159
|
+
chat_template_kwargs={"enable_thinking": False},
|
|
160
|
+
)
|
|
161
|
+
assert result == Sentiment(label="negative", confidence=0.95)
|
|
162
|
+
body = svc._session.post.call_args.kwargs["json"]
|
|
163
|
+
assert body["response_format"] == response_format_for(Sentiment)
|
|
164
|
+
assert body["max_tokens"] == 2000
|
|
165
|
+
assert body["chat_template_kwargs"] == {"enable_thinking": False}
|
|
166
|
+
|
|
167
|
+
|
|
168
|
+
def test_completions_structured_refuses_streaming():
|
|
169
|
+
svc = _sync_service(_response(GOOD))
|
|
170
|
+
with pytest.raises(TypeError):
|
|
171
|
+
svc.completions_structured("m", [], Sentiment, stream=True)
|
|
172
|
+
svc._session.post.assert_not_called()
|
|
173
|
+
|
|
174
|
+
|
|
175
|
+
def test_async_completions_structured_sends_schema_and_returns_instance():
|
|
176
|
+
client = FakeAsyncClient(
|
|
177
|
+
responses=[
|
|
178
|
+
FakeAsyncResponse(200, headers=rl_headers(), json_body=_response(GOOD))
|
|
179
|
+
]
|
|
180
|
+
)
|
|
181
|
+
svc = AsyncChatService(client, "https://x/v1")
|
|
182
|
+
result = asyncio.run(svc.completions_structured("m", [], Sentiment))
|
|
183
|
+
assert result == Sentiment(label="negative", confidence=0.95)
|
|
184
|
+
assert client.calls[0]["json"]["response_format"] == response_format_for(Sentiment)
|
|
@@ -1,81 +0,0 @@
|
|
|
1
|
-
"""Chat service — completions and streaming."""
|
|
2
|
-
|
|
3
|
-
from __future__ import annotations
|
|
4
|
-
|
|
5
|
-
from typing import TYPE_CHECKING
|
|
6
|
-
|
|
7
|
-
from ._http import RetryPolicy, coerce_retry, post_chat_completion, resolve_retry
|
|
8
|
-
from ._streaming import SSEStream
|
|
9
|
-
|
|
10
|
-
if TYPE_CHECKING:
|
|
11
|
-
import requests
|
|
12
|
-
|
|
13
|
-
|
|
14
|
-
class ChatService:
|
|
15
|
-
"""Access the ``/chat/completions`` endpoint.
|
|
16
|
-
|
|
17
|
-
Args:
|
|
18
|
-
session: A :class:`requests.Session` with auth headers configured.
|
|
19
|
-
base_url: The SAIA API base URL.
|
|
20
|
-
"""
|
|
21
|
-
|
|
22
|
-
def __init__(
|
|
23
|
-
self,
|
|
24
|
-
session: requests.Session,
|
|
25
|
-
base_url: str,
|
|
26
|
-
*,
|
|
27
|
-
retry: RetryPolicy | bool | None = None,
|
|
28
|
-
):
|
|
29
|
-
self._session = session
|
|
30
|
-
self._base_url = base_url
|
|
31
|
-
self._retry = coerce_retry(retry)
|
|
32
|
-
|
|
33
|
-
def completions(
|
|
34
|
-
self,
|
|
35
|
-
model: str,
|
|
36
|
-
messages: list[dict],
|
|
37
|
-
*,
|
|
38
|
-
temperature: float | None = None,
|
|
39
|
-
top_p: float | None = None,
|
|
40
|
-
max_tokens: int | None = None,
|
|
41
|
-
stream: bool = False,
|
|
42
|
-
retry: RetryPolicy | bool | None = None,
|
|
43
|
-
**kwargs,
|
|
44
|
-
) -> dict | SSEStream:
|
|
45
|
-
"""Send a chat completion request.
|
|
46
|
-
|
|
47
|
-
Args:
|
|
48
|
-
model: Model identifier (e.g. ``"meta-llama-3.1-8b-instruct"``).
|
|
49
|
-
messages: List of message dicts with ``"role"`` and ``"content"`` keys.
|
|
50
|
-
temperature: Sampling temperature (0–2).
|
|
51
|
-
top_p: Nucleus sampling parameter (0–1).
|
|
52
|
-
max_tokens: Maximum tokens to generate.
|
|
53
|
-
stream: If ``True``, return a generator yielding chunks.
|
|
54
|
-
**kwargs: Additional parameters forwarded to the API.
|
|
55
|
-
|
|
56
|
-
Returns:
|
|
57
|
-
When ``stream=False``: the API response dict, with an extra
|
|
58
|
-
``"_rate_limits"`` key — a JSON-serializable dict of the current
|
|
59
|
-
rate-limit headers (see :class:`~saia_python.RateLimitInfo`).
|
|
60
|
-
When ``stream=True``: an ``SSEStream`` — iterate it for the
|
|
61
|
-
response chunks; its ``rate_limits`` attribute exposes the same
|
|
62
|
-
dict (available immediately, from the response headers).
|
|
63
|
-
"""
|
|
64
|
-
body = {"model": model, "messages": messages, **kwargs}
|
|
65
|
-
if temperature is not None:
|
|
66
|
-
body["temperature"] = temperature
|
|
67
|
-
if top_p is not None:
|
|
68
|
-
body["top_p"] = top_p
|
|
69
|
-
if max_tokens is not None:
|
|
70
|
-
body["max_tokens"] = max_tokens
|
|
71
|
-
|
|
72
|
-
return post_chat_completion(
|
|
73
|
-
self._session,
|
|
74
|
-
f"{self._base_url}/chat/completions",
|
|
75
|
-
body,
|
|
76
|
-
stream=stream,
|
|
77
|
-
policy=resolve_retry(self._retry, retry),
|
|
78
|
-
)
|
|
79
|
-
|
|
80
|
-
def __repr__(self):
|
|
81
|
-
return f"ChatService(base_url={self._base_url!r})"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|