sthai 0.1.0__py3-none-any.whl
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.
- sthai/__init__.py +12 -0
- sthai/client.py +563 -0
- sthai/const.py +52 -0
- sthai/models.py +13 -0
- sthai/py.typed +0 -0
- sthai/structs/__init__.py +101 -0
- sthai/structs/common.py +23 -0
- sthai/structs/completions.py +303 -0
- sthai/structs/embeddings.py +127 -0
- sthai/structs/models.py +39 -0
- sthai/structs/rerank.py +95 -0
- sthai/typing.py +22 -0
- sthai-0.1.0.dist-info/METADATA +211 -0
- sthai-0.1.0.dist-info/RECORD +16 -0
- sthai-0.1.0.dist-info/WHEEL +4 -0
- sthai-0.1.0.dist-info/licenses/LICENSE +21 -0
sthai/__init__.py
ADDED
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
"""Python client for the SiteHost AI Platform (inference, embeddings, reranking)."""
|
|
2
|
+
|
|
3
|
+
from sthai.client import Client, image_content
|
|
4
|
+
from sthai.models import EmbeddingModel, InferenceModel, RerankingModel
|
|
5
|
+
|
|
6
|
+
__all__ = [
|
|
7
|
+
"Client",
|
|
8
|
+
"EmbeddingModel",
|
|
9
|
+
"InferenceModel",
|
|
10
|
+
"RerankingModel",
|
|
11
|
+
"image_content",
|
|
12
|
+
]
|
sthai/client.py
ADDED
|
@@ -0,0 +1,563 @@
|
|
|
1
|
+
import warnings
|
|
2
|
+
from base64 import b64encode
|
|
3
|
+
from collections.abc import Sequence
|
|
4
|
+
from os import getenv
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
from secrets import token_hex
|
|
7
|
+
from typing import Any, TypeVar
|
|
8
|
+
|
|
9
|
+
import msgspec
|
|
10
|
+
from msgspec import UNSET
|
|
11
|
+
from niquests import request
|
|
12
|
+
from niquests.models import Response
|
|
13
|
+
|
|
14
|
+
from sthai.const import (
|
|
15
|
+
EMBEDDING_ENDPOINT,
|
|
16
|
+
EMBEDDING_PARAMS,
|
|
17
|
+
HEALTH_ENDPOINT,
|
|
18
|
+
INFERENCE_ENDPOINT,
|
|
19
|
+
MODELS_ENDPOINT,
|
|
20
|
+
RERANKING_ENDPOINT,
|
|
21
|
+
SESSION_PIN_HEADER,
|
|
22
|
+
)
|
|
23
|
+
from sthai.models import EmbeddingModel, InferenceModel, RerankingModel
|
|
24
|
+
from sthai.structs.completions import (
|
|
25
|
+
AssistantMessage,
|
|
26
|
+
ChatMessage,
|
|
27
|
+
ContentPart,
|
|
28
|
+
ImageContent,
|
|
29
|
+
ImageURL,
|
|
30
|
+
InferenceRequest,
|
|
31
|
+
InferenceResponse,
|
|
32
|
+
JsonSchemaResponseFormat,
|
|
33
|
+
ResponseFormat,
|
|
34
|
+
SystemMessage,
|
|
35
|
+
TextContent,
|
|
36
|
+
UserMessage,
|
|
37
|
+
)
|
|
38
|
+
from sthai.structs.embeddings import (
|
|
39
|
+
EmbeddingChatRequest,
|
|
40
|
+
EmbeddingCompletionRequest,
|
|
41
|
+
EmbeddingResponse,
|
|
42
|
+
)
|
|
43
|
+
from sthai.structs.models import ModelCard, ModelList
|
|
44
|
+
from sthai.structs.rerank import (
|
|
45
|
+
RerankRequest,
|
|
46
|
+
RerankResponse,
|
|
47
|
+
RerankResult,
|
|
48
|
+
ScoreMultiModalParam,
|
|
49
|
+
)
|
|
50
|
+
from sthai.typing import HttpMethod
|
|
51
|
+
|
|
52
|
+
# TypeVar rather than PEP 695 syntax to stay compatible with Python 3.10
|
|
53
|
+
T = TypeVar("T")
|
|
54
|
+
|
|
55
|
+
# Magic-byte prefixes for the image formats the API accepts
|
|
56
|
+
_IMAGE_MAGIC_BYTES = {
|
|
57
|
+
b"\x89PNG": "image/png",
|
|
58
|
+
b"\xff\xd8\xff": "image/jpeg",
|
|
59
|
+
b"GIF8": "image/gif",
|
|
60
|
+
}
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def _image_file_to_data_uri(image: Path | bytes) -> str:
|
|
64
|
+
"""
|
|
65
|
+
Base64-encode an image file (or raw image bytes) into a data URI
|
|
66
|
+
suitable for the image_url content part of a chat message.
|
|
67
|
+
"""
|
|
68
|
+
data = image.read_bytes() if isinstance(image, Path) else image
|
|
69
|
+
if data[:4] == b"RIFF" and data[8:12] == b"WEBP":
|
|
70
|
+
mime = "image/webp"
|
|
71
|
+
else:
|
|
72
|
+
for magic, magic_mime in _IMAGE_MAGIC_BYTES.items():
|
|
73
|
+
if data.startswith(magic):
|
|
74
|
+
mime = magic_mime
|
|
75
|
+
break
|
|
76
|
+
else:
|
|
77
|
+
raise ValueError(
|
|
78
|
+
"unrecognized image format: expected PNG, JPEG, GIF, or WEBP"
|
|
79
|
+
)
|
|
80
|
+
return f"data:{mime};base64,{b64encode(data).decode('ascii')}"
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
def image_content(image: str | Path | bytes) -> ImageContent:
|
|
84
|
+
"""
|
|
85
|
+
Build an image_url content part from a URL string, a local file path,
|
|
86
|
+
or raw image bytes, for hand-constructing multimodal content.
|
|
87
|
+
"""
|
|
88
|
+
if isinstance(image, str):
|
|
89
|
+
return ImageContent(image_url=ImageURL(url=image))
|
|
90
|
+
return ImageContent(image_url=ImageURL(url=_image_file_to_data_uri(image)))
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def _build_image_parts(
|
|
94
|
+
image_urls: list[str] | None,
|
|
95
|
+
image_files: list[Path | bytes] | None,
|
|
96
|
+
) -> list[ImageContent]:
|
|
97
|
+
"""Image content parts for the given URLs and local files or bytes."""
|
|
98
|
+
return [
|
|
99
|
+
image_content(image) for image in [*(image_urls or []), *(image_files or [])]
|
|
100
|
+
]
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
def _prompt_content(
|
|
104
|
+
prompt: str,
|
|
105
|
+
image_urls: list[str] | None,
|
|
106
|
+
image_files: list[Path | bytes] | None,
|
|
107
|
+
) -> str | list[ContentPart]:
|
|
108
|
+
"""A user-message content value: plain text, or text plus image parts."""
|
|
109
|
+
image_parts = _build_image_parts(image_urls, image_files)
|
|
110
|
+
return [TextContent(text=prompt), *image_parts] if image_parts else prompt
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
def _response_format(response_type: type) -> ResponseFormat:
|
|
114
|
+
"""
|
|
115
|
+
The response_format for a structured response: plain JSON mode for dict,
|
|
116
|
+
otherwise a JSON schema generated from the type for the server to
|
|
117
|
+
enforce via guided decoding.
|
|
118
|
+
"""
|
|
119
|
+
if response_type is dict:
|
|
120
|
+
return ResponseFormat(type="json_object")
|
|
121
|
+
return ResponseFormat(
|
|
122
|
+
type="json_schema",
|
|
123
|
+
json_schema=JsonSchemaResponseFormat(
|
|
124
|
+
name=getattr(response_type, "__name__", "response"),
|
|
125
|
+
json_schema=msgspec.json.schema(response_type),
|
|
126
|
+
),
|
|
127
|
+
)
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
def _default_instruction(model: EmbeddingModel | str, query: bool) -> str | None:
|
|
131
|
+
"""
|
|
132
|
+
The model's recommended embedding instruction, if known: its query
|
|
133
|
+
instruction when query is set, its document instruction otherwise.
|
|
134
|
+
"""
|
|
135
|
+
params = EMBEDDING_PARAMS.get(model)
|
|
136
|
+
if params is None:
|
|
137
|
+
return None
|
|
138
|
+
return params.query_instruction if query else params.document_instruction
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
def _check_dimensions(model: EmbeddingModel | str, dimensions: int | None) -> None:
|
|
142
|
+
"""
|
|
143
|
+
Warn when a requested Matryoshka truncation exceeds or does not divide
|
|
144
|
+
evenly into the model's native output dimension (when both are known).
|
|
145
|
+
"""
|
|
146
|
+
if dimensions is None:
|
|
147
|
+
return
|
|
148
|
+
if dimensions < 1:
|
|
149
|
+
raise ValueError("dimensions must be a positive integer")
|
|
150
|
+
params = EMBEDDING_PARAMS.get(model)
|
|
151
|
+
if params is None or params.dimensions is None:
|
|
152
|
+
return
|
|
153
|
+
if dimensions > params.dimensions:
|
|
154
|
+
warnings.warn(
|
|
155
|
+
f"dimensions={dimensions} exceeds the native {params.dimensions} "
|
|
156
|
+
f"dimensions of '{model}'",
|
|
157
|
+
stacklevel=3,
|
|
158
|
+
)
|
|
159
|
+
elif params.dimensions % dimensions != 0:
|
|
160
|
+
warnings.warn(
|
|
161
|
+
f"dimensions={dimensions} does not divide evenly into the native "
|
|
162
|
+
f"{params.dimensions} dimensions of '{model}'; use a power-of-two "
|
|
163
|
+
f"divisor (e.g. {params.dimensions // 2}, {params.dimensions // 4})",
|
|
164
|
+
stacklevel=3,
|
|
165
|
+
)
|
|
166
|
+
|
|
167
|
+
|
|
168
|
+
def _float_embedding(embedding: list[float] | str) -> list[float]:
|
|
169
|
+
"""
|
|
170
|
+
Ensure a decoded embedding is the float list the client requested (the
|
|
171
|
+
server returns strings for non-float encoding formats).
|
|
172
|
+
"""
|
|
173
|
+
if isinstance(embedding, str):
|
|
174
|
+
raise TypeError("expected a float embedding, got an encoded string")
|
|
175
|
+
return embedding
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
class Client:
|
|
179
|
+
"""
|
|
180
|
+
The main client class for interacting with the SthAI API.
|
|
181
|
+
"""
|
|
182
|
+
|
|
183
|
+
def __init__(
|
|
184
|
+
self,
|
|
185
|
+
api_key: str = getenv("STHAI_KEY", ""),
|
|
186
|
+
*,
|
|
187
|
+
fqdn: str = "ai.sitehost.nz",
|
|
188
|
+
secure: bool = True,
|
|
189
|
+
session_pin: str | None = None,
|
|
190
|
+
auto_session: bool = False,
|
|
191
|
+
write_history: bool = True,
|
|
192
|
+
) -> None:
|
|
193
|
+
"""
|
|
194
|
+
Create a client. api_key defaults to the STHAI_KEY environment
|
|
195
|
+
variable. session_pin (or auto_session=True to generate one) pins
|
|
196
|
+
requests to a server session; write_history controls whether chat()
|
|
197
|
+
records conversation turns.
|
|
198
|
+
"""
|
|
199
|
+
self.fqdn = fqdn
|
|
200
|
+
self.secure = secure
|
|
201
|
+
if not api_key:
|
|
202
|
+
raise ValueError("api_key is required")
|
|
203
|
+
self._api_key = api_key
|
|
204
|
+
self._session_pin = session_pin
|
|
205
|
+
if not self._session_pin and auto_session:
|
|
206
|
+
self._session_pin = token_hex(24)
|
|
207
|
+
self._write_history = write_history
|
|
208
|
+
self._chat_history: list[ChatMessage] = []
|
|
209
|
+
self._last_response: InferenceResponse | None = None
|
|
210
|
+
|
|
211
|
+
def healthy(self) -> bool:
|
|
212
|
+
"""Check the server's health status."""
|
|
213
|
+
response = self._make_request(HttpMethod.GET, HEALTH_ENDPOINT)
|
|
214
|
+
return response.ok
|
|
215
|
+
|
|
216
|
+
def models(self) -> list[ModelCard]:
|
|
217
|
+
"""List the models available on the server."""
|
|
218
|
+
response = self._make_request(HttpMethod.GET, MODELS_ENDPOINT)
|
|
219
|
+
response.raise_for_status()
|
|
220
|
+
return msgspec.json.decode(response.content or b"", type=ModelList).data
|
|
221
|
+
|
|
222
|
+
def clear_history(self) -> None:
|
|
223
|
+
"""Discard the stored chat history."""
|
|
224
|
+
self._chat_history = []
|
|
225
|
+
|
|
226
|
+
def last_response(self) -> InferenceResponse | None:
|
|
227
|
+
"""The full response from the most recent chat() or response() call."""
|
|
228
|
+
return self._last_response
|
|
229
|
+
|
|
230
|
+
def last_reasoning(self) -> str | None:
|
|
231
|
+
"""The reasoning from the most recent chat() or response() call, if any."""
|
|
232
|
+
if self._last_response is None:
|
|
233
|
+
return None
|
|
234
|
+
return self._last_response.output().reasoning
|
|
235
|
+
|
|
236
|
+
def chat(
|
|
237
|
+
self,
|
|
238
|
+
prompt: str,
|
|
239
|
+
*,
|
|
240
|
+
model: InferenceModel | str = InferenceModel.QWEN_3_6_27B,
|
|
241
|
+
max_tokens: int | None = None,
|
|
242
|
+
temperature: float | None = None,
|
|
243
|
+
use_thinking: bool = False,
|
|
244
|
+
system_prompt: str | None = None,
|
|
245
|
+
image_urls: list[str] | None = None,
|
|
246
|
+
image_files: list[Path | bytes] | None = None,
|
|
247
|
+
use_history: bool = True,
|
|
248
|
+
) -> InferenceResponse:
|
|
249
|
+
"""
|
|
250
|
+
Send a chat message and return the full inference response.
|
|
251
|
+
|
|
252
|
+
With write_history=True (the default), each successful call appends
|
|
253
|
+
the user and assistant turns to the stored history, and later calls
|
|
254
|
+
send that history. Pass use_history=False for a standalone call that
|
|
255
|
+
neither sends nor updates it.
|
|
256
|
+
"""
|
|
257
|
+
user_message = UserMessage(
|
|
258
|
+
content=_prompt_content(prompt, image_urls, image_files)
|
|
259
|
+
)
|
|
260
|
+
|
|
261
|
+
messages: list[ChatMessage] = []
|
|
262
|
+
if system_prompt:
|
|
263
|
+
# The system prompt is prepended per-call rather than stored in
|
|
264
|
+
# history, so changing it between calls behaves predictably
|
|
265
|
+
messages.append(SystemMessage(content=system_prompt))
|
|
266
|
+
if use_history:
|
|
267
|
+
messages.extend(self._chat_history)
|
|
268
|
+
messages.append(user_message)
|
|
269
|
+
|
|
270
|
+
body = InferenceRequest(
|
|
271
|
+
messages=messages,
|
|
272
|
+
model=model,
|
|
273
|
+
# max_tokens is deprecated upstream in favor of max_completion_tokens
|
|
274
|
+
max_completion_tokens=max_tokens if max_tokens is not None else UNSET,
|
|
275
|
+
temperature=temperature if temperature is not None else UNSET,
|
|
276
|
+
chat_template_kwargs={"enable_thinking": use_thinking},
|
|
277
|
+
)
|
|
278
|
+
decoded = self._inference_request(body)
|
|
279
|
+
|
|
280
|
+
# A call that didn't see the history must not write to it either,
|
|
281
|
+
# or the stored conversation would gain a turn with missing context
|
|
282
|
+
if use_history and self._write_history and decoded.choices:
|
|
283
|
+
self._chat_history.append(user_message)
|
|
284
|
+
self._chat_history.append(
|
|
285
|
+
AssistantMessage(content=decoded.choices[0].message.content)
|
|
286
|
+
)
|
|
287
|
+
return decoded
|
|
288
|
+
|
|
289
|
+
def response(
|
|
290
|
+
self,
|
|
291
|
+
prompt: str,
|
|
292
|
+
*,
|
|
293
|
+
response_type: type[T] | None = None,
|
|
294
|
+
model: InferenceModel | str = InferenceModel.QWEN_3_6_27B,
|
|
295
|
+
max_tokens: int | None = None,
|
|
296
|
+
temperature: float | None = None,
|
|
297
|
+
use_thinking: bool = False,
|
|
298
|
+
system_prompt: str | None = None,
|
|
299
|
+
image_urls: list[str] | None = None,
|
|
300
|
+
image_files: list[Path | bytes] | None = None,
|
|
301
|
+
) -> T | InferenceResponse:
|
|
302
|
+
"""
|
|
303
|
+
One-off inference: like chat(), but the stored chat history is
|
|
304
|
+
neither sent nor updated. The response remains available through
|
|
305
|
+
last_response() and last_reasoning().
|
|
306
|
+
|
|
307
|
+
Pass response_type for a structured response: a msgspec Struct type
|
|
308
|
+
becomes a JSON schema the server enforces during generation, and the
|
|
309
|
+
response text is decoded and validated into that type and returned.
|
|
310
|
+
response_type=dict just constrains the output to valid JSON. Without
|
|
311
|
+
response_type the full InferenceResponse is returned, as with chat().
|
|
312
|
+
|
|
313
|
+
Thinking combines with structured responses (reasoning stays
|
|
314
|
+
unconstrained) but consumes max_tokens, so budget generously. The
|
|
315
|
+
server occasionally skips the schema when thinking is enabled;
|
|
316
|
+
parsing then raises a ValueError - retry, or disable thinking.
|
|
317
|
+
"""
|
|
318
|
+
messages: list[ChatMessage] = []
|
|
319
|
+
if system_prompt:
|
|
320
|
+
messages.append(SystemMessage(content=system_prompt))
|
|
321
|
+
messages.append(
|
|
322
|
+
UserMessage(content=_prompt_content(prompt, image_urls, image_files))
|
|
323
|
+
)
|
|
324
|
+
|
|
325
|
+
body = InferenceRequest(
|
|
326
|
+
messages=messages,
|
|
327
|
+
model=model,
|
|
328
|
+
max_completion_tokens=max_tokens if max_tokens is not None else UNSET,
|
|
329
|
+
temperature=temperature if temperature is not None else UNSET,
|
|
330
|
+
chat_template_kwargs={"enable_thinking": use_thinking},
|
|
331
|
+
response_format=(
|
|
332
|
+
_response_format(response_type) if response_type is not None else UNSET
|
|
333
|
+
),
|
|
334
|
+
)
|
|
335
|
+
decoded = self._inference_request(body)
|
|
336
|
+
if response_type is None:
|
|
337
|
+
return decoded
|
|
338
|
+
return decoded.parse(response_type)
|
|
339
|
+
|
|
340
|
+
def _inference_request(self, body: InferenceRequest) -> InferenceResponse:
|
|
341
|
+
"""POST an inference request and record the decoded last response."""
|
|
342
|
+
response = self._make_request(
|
|
343
|
+
HttpMethod.POST,
|
|
344
|
+
INFERENCE_ENDPOINT,
|
|
345
|
+
data=msgspec.json.encode(body),
|
|
346
|
+
headers={"Content-Type": "application/json"},
|
|
347
|
+
)
|
|
348
|
+
response.raise_for_status()
|
|
349
|
+
decoded = msgspec.json.decode(response.content or b"", type=InferenceResponse)
|
|
350
|
+
self._last_response = decoded
|
|
351
|
+
return decoded
|
|
352
|
+
|
|
353
|
+
def embed(
|
|
354
|
+
self,
|
|
355
|
+
text: str | None = None,
|
|
356
|
+
*,
|
|
357
|
+
model: EmbeddingModel | str = EmbeddingModel.QWEN_3_VL_8B,
|
|
358
|
+
query: bool = False,
|
|
359
|
+
instruction: str | None = None,
|
|
360
|
+
image_urls: list[str] | None = None,
|
|
361
|
+
image_files: list[Path | bytes] | None = None,
|
|
362
|
+
dimensions: int | None = None,
|
|
363
|
+
) -> list[float]:
|
|
364
|
+
"""
|
|
365
|
+
Embed a single input (text, images, or both) and return its vector.
|
|
366
|
+
|
|
367
|
+
The instruction-trained model is steered by a default instruction
|
|
368
|
+
from EMBEDDING_PARAMS in sthai.const: the model's document
|
|
369
|
+
instruction, or its query instruction when query=True (use this when
|
|
370
|
+
embedding search queries). Passing instruction overrides either.
|
|
371
|
+
|
|
372
|
+
Each call produces exactly ONE vector - multimodal content rolls into
|
|
373
|
+
it; use batch_embed() to embed many texts in one request. dimensions
|
|
374
|
+
truncates the vector server-side (Matryoshka); powers of two work
|
|
375
|
+
best, up to the model's native dimension.
|
|
376
|
+
"""
|
|
377
|
+
image_parts = _build_image_parts(image_urls, image_files)
|
|
378
|
+
parts: list[ContentPart] = []
|
|
379
|
+
# An empty string is treated as no text: embedding it would produce a
|
|
380
|
+
# meaningless vector, so it falls through to the guard below instead
|
|
381
|
+
if text:
|
|
382
|
+
parts.append(TextContent(text=text))
|
|
383
|
+
parts.extend(image_parts)
|
|
384
|
+
if not parts:
|
|
385
|
+
raise ValueError("embed() requires text and/or images")
|
|
386
|
+
_check_dimensions(model, dimensions)
|
|
387
|
+
content: str | list[ContentPart] = text if text and not image_parts else parts
|
|
388
|
+
|
|
389
|
+
if instruction is None:
|
|
390
|
+
instruction = _default_instruction(model, query)
|
|
391
|
+
messages: list[ChatMessage] = []
|
|
392
|
+
if instruction is not None:
|
|
393
|
+
messages.append(SystemMessage(content=instruction))
|
|
394
|
+
messages.append(UserMessage(content=content))
|
|
395
|
+
# The open assistant turn is intentional: with continue_final_message
|
|
396
|
+
# the template is left unterminated, matching how the model was
|
|
397
|
+
# trained to embed
|
|
398
|
+
messages.append(AssistantMessage(content=""))
|
|
399
|
+
|
|
400
|
+
body = EmbeddingChatRequest(
|
|
401
|
+
messages=messages,
|
|
402
|
+
model=model,
|
|
403
|
+
encoding_format="float",
|
|
404
|
+
dimensions=dimensions if dimensions is not None else UNSET,
|
|
405
|
+
continue_final_message=True,
|
|
406
|
+
# True (not the chat-form server default of False) so tokenization
|
|
407
|
+
# matches batch_embed's plain-input form, which defaults to True
|
|
408
|
+
add_special_tokens=True,
|
|
409
|
+
)
|
|
410
|
+
decoded = self._embedding_request(body)
|
|
411
|
+
outputs = decoded.output()
|
|
412
|
+
if not outputs:
|
|
413
|
+
raise ValueError("server returned no embedding data")
|
|
414
|
+
return _float_embedding(outputs[0])
|
|
415
|
+
|
|
416
|
+
def batch_embed(
|
|
417
|
+
self,
|
|
418
|
+
texts: list[str],
|
|
419
|
+
*,
|
|
420
|
+
model: EmbeddingModel | str = EmbeddingModel.QWEN_3_VL_8B,
|
|
421
|
+
query: bool = False,
|
|
422
|
+
instruction: str | None = None,
|
|
423
|
+
template: str | None = None,
|
|
424
|
+
dimensions: int | None = None,
|
|
425
|
+
) -> list[list[float]]:
|
|
426
|
+
"""
|
|
427
|
+
Embed a batch of texts in one request, returning one vector per text
|
|
428
|
+
in the same order. Text-only; use embed() for multimodal input.
|
|
429
|
+
|
|
430
|
+
Only the plain-input request form batches, and it bypasses the
|
|
431
|
+
server-side chat template, so each text is rendered through a local
|
|
432
|
+
template first; with the built-in templates the results match calling
|
|
433
|
+
embed() per text. template and instruction default from
|
|
434
|
+
EMBEDDING_PARAMS (the query instruction when query=True, as with
|
|
435
|
+
embed()). For models without known params, pass a template using
|
|
436
|
+
{instruction} and {text} placeholders - "{text}" alone for raw
|
|
437
|
+
untemplated input. dimensions truncates the vectors server-side.
|
|
438
|
+
"""
|
|
439
|
+
if not texts:
|
|
440
|
+
raise ValueError("batch_embed() requires at least one text")
|
|
441
|
+
if not all(texts):
|
|
442
|
+
raise ValueError("batch_embed() texts must be non-empty strings")
|
|
443
|
+
if template is None:
|
|
444
|
+
params = EMBEDDING_PARAMS.get(model)
|
|
445
|
+
template = params.template if params is not None else None
|
|
446
|
+
if template is None:
|
|
447
|
+
raise ValueError(
|
|
448
|
+
f"no known embedding template for model '{model}'; pass "
|
|
449
|
+
'template= (use "{text}" for models that take raw '
|
|
450
|
+
"untemplated input)"
|
|
451
|
+
)
|
|
452
|
+
if "{instruction}" not in template and (instruction is not None or query):
|
|
453
|
+
warnings.warn(
|
|
454
|
+
"the template has no {instruction} placeholder, so the requested "
|
|
455
|
+
"instruction steering will not be applied",
|
|
456
|
+
stacklevel=2,
|
|
457
|
+
)
|
|
458
|
+
if instruction is None:
|
|
459
|
+
instruction = _default_instruction(model, query)
|
|
460
|
+
if instruction is None and "{instruction}" in template:
|
|
461
|
+
raise ValueError(
|
|
462
|
+
f"no known embedding instruction for model '{model}' but "
|
|
463
|
+
"the template expects one; pass instruction="
|
|
464
|
+
)
|
|
465
|
+
_check_dimensions(model, dimensions)
|
|
466
|
+
try:
|
|
467
|
+
inputs = [
|
|
468
|
+
template.format(instruction=instruction, text=text) for text in texts
|
|
469
|
+
]
|
|
470
|
+
except (KeyError, IndexError) as exc:
|
|
471
|
+
raise ValueError(
|
|
472
|
+
"template must use only the {instruction} and {text} "
|
|
473
|
+
"placeholders; escape literal braces as {{ and }}"
|
|
474
|
+
) from exc
|
|
475
|
+
body = EmbeddingCompletionRequest(
|
|
476
|
+
input=inputs,
|
|
477
|
+
model=model,
|
|
478
|
+
encoding_format="float",
|
|
479
|
+
dimensions=dimensions if dimensions is not None else UNSET,
|
|
480
|
+
)
|
|
481
|
+
decoded = self._embedding_request(body)
|
|
482
|
+
return [_float_embedding(embedding) for embedding in decoded.output()]
|
|
483
|
+
|
|
484
|
+
def _embedding_request(
|
|
485
|
+
self, body: EmbeddingChatRequest | EmbeddingCompletionRequest
|
|
486
|
+
) -> EmbeddingResponse:
|
|
487
|
+
"""POST an embedding request and decode the response."""
|
|
488
|
+
response = self._make_request(
|
|
489
|
+
HttpMethod.POST,
|
|
490
|
+
EMBEDDING_ENDPOINT,
|
|
491
|
+
data=msgspec.json.encode(body),
|
|
492
|
+
headers={"Content-Type": "application/json"},
|
|
493
|
+
)
|
|
494
|
+
response.raise_for_status()
|
|
495
|
+
return msgspec.json.decode(response.content or b"", type=EmbeddingResponse)
|
|
496
|
+
|
|
497
|
+
def rerank(
|
|
498
|
+
self,
|
|
499
|
+
query: str | ScoreMultiModalParam,
|
|
500
|
+
# Sequence rather than list so a plain list[str] type-checks
|
|
501
|
+
documents: Sequence[str | ScoreMultiModalParam],
|
|
502
|
+
*,
|
|
503
|
+
model: RerankingModel | str = RerankingModel.QWEN_3_VL_8B,
|
|
504
|
+
top_n: int | None = None,
|
|
505
|
+
instruction: str | None = None,
|
|
506
|
+
) -> list[RerankResult]:
|
|
507
|
+
"""
|
|
508
|
+
Score each document against the query and return the results sorted
|
|
509
|
+
by relevance score descending, each carrying the document, its
|
|
510
|
+
relevance_score, and its index in the input documents list.
|
|
511
|
+
|
|
512
|
+
All documents are returned unless top_n limits it. The
|
|
513
|
+
instruction-trained model applies its own default instruction; pass
|
|
514
|
+
instruction to steer relevance for a specific task. The query and
|
|
515
|
+
each document may be a plain string or, for multimodal input, a
|
|
516
|
+
ScoreMultiModalParam wrapping text/image content parts.
|
|
517
|
+
"""
|
|
518
|
+
if not documents:
|
|
519
|
+
raise ValueError("rerank() requires at least one document")
|
|
520
|
+
if top_n is not None and top_n < 1:
|
|
521
|
+
raise ValueError("top_n must be a positive integer")
|
|
522
|
+
|
|
523
|
+
body = RerankRequest(
|
|
524
|
+
query=query,
|
|
525
|
+
documents=list(documents),
|
|
526
|
+
top_n=top_n if top_n is not None else UNSET,
|
|
527
|
+
model=model,
|
|
528
|
+
instruction=instruction if instruction is not None else UNSET,
|
|
529
|
+
)
|
|
530
|
+
response = self._make_request(
|
|
531
|
+
HttpMethod.POST,
|
|
532
|
+
RERANKING_ENDPOINT,
|
|
533
|
+
data=msgspec.json.encode(body),
|
|
534
|
+
headers={"Content-Type": "application/json"},
|
|
535
|
+
)
|
|
536
|
+
response.raise_for_status()
|
|
537
|
+
decoded = msgspec.json.decode(response.content or b"", type=RerankResponse)
|
|
538
|
+
return decoded.results
|
|
539
|
+
|
|
540
|
+
def _server_url(self) -> str:
|
|
541
|
+
"""The server's base URL."""
|
|
542
|
+
scheme = "https" if self.secure else "http"
|
|
543
|
+
return f"{scheme}://{self.fqdn}"
|
|
544
|
+
|
|
545
|
+
def _endpoint_url(self, endpoint: str) -> str:
|
|
546
|
+
"""The absolute URL for an endpoint path."""
|
|
547
|
+
if not endpoint.startswith("/"):
|
|
548
|
+
endpoint = f"/{endpoint}"
|
|
549
|
+
return f"{self._server_url()}{endpoint}"
|
|
550
|
+
|
|
551
|
+
def _default_headers(self) -> dict[str, str]:
|
|
552
|
+
"""The auth (and session pin) headers sent with every request."""
|
|
553
|
+
headers = {"Authorization": f"Bearer {self._api_key}"}
|
|
554
|
+
if self._session_pin:
|
|
555
|
+
headers[SESSION_PIN_HEADER] = self._session_pin
|
|
556
|
+
return headers
|
|
557
|
+
|
|
558
|
+
def _make_request(
|
|
559
|
+
self, method: HttpMethod, endpoint: str, **kwargs: Any
|
|
560
|
+
) -> Response:
|
|
561
|
+
"""Send a request with default headers; kwargs pass through to niquests."""
|
|
562
|
+
headers = self._default_headers() | kwargs.pop("headers", {})
|
|
563
|
+
return request(method, self._endpoint_url(endpoint), headers=headers, **kwargs)
|
sthai/const.py
ADDED
|
@@ -0,0 +1,52 @@
|
|
|
1
|
+
from typing import NamedTuple
|
|
2
|
+
|
|
3
|
+
from sthai.models import EmbeddingModel
|
|
4
|
+
|
|
5
|
+
INFERENCE_ENDPOINT = "/v1/chat/completions"
|
|
6
|
+
EMBEDDING_ENDPOINT = "/v1/embeddings"
|
|
7
|
+
RERANKING_ENDPOINT = "/v1/rerank"
|
|
8
|
+
MODELS_ENDPOINT = "/v1/models"
|
|
9
|
+
HEALTH_ENDPOINT = "/health"
|
|
10
|
+
|
|
11
|
+
SESSION_PIN_HEADER = "X-Session-Id"
|
|
12
|
+
|
|
13
|
+
# Chat templates applied locally for batched embedding requests, since only
|
|
14
|
+
# the plain-input request form batches and it bypasses the server-side
|
|
15
|
+
# template. Each matches what the model's own chat template renders for a
|
|
16
|
+
# single-turn request; the open assistant turn is intentional.
|
|
17
|
+
QWEN_3_VL_EMBEDDING_TEMPLATE = (
|
|
18
|
+
"<|im_start|>system\n{instruction}<|im_end|>\n"
|
|
19
|
+
"<|im_start|>user\n{text}<|im_end|>\n"
|
|
20
|
+
"<|im_start|>assistant\n"
|
|
21
|
+
)
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
# Recommended task instructions for the embedding model
|
|
25
|
+
EMBEDDING_DOCUMENT_INSTRUCTION = "Represent the user's input."
|
|
26
|
+
EMBEDDING_QUERY_INSTRUCTION = (
|
|
27
|
+
"Given a web search query, retrieve relevant passages that answer the query"
|
|
28
|
+
)
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
class EmbeddingParams(NamedTuple):
|
|
32
|
+
"""Client-side defaults for a known embedding model."""
|
|
33
|
+
|
|
34
|
+
# Chat template applied locally for batched requests (see batch_embed)
|
|
35
|
+
template: str | None = None
|
|
36
|
+
# The model's native output dimension; a requested Matryoshka truncation
|
|
37
|
+
# should divide evenly into this
|
|
38
|
+
dimensions: int | None = None
|
|
39
|
+
# Recommended task instructions: the default when embedding documents,
|
|
40
|
+
# and the one to pass explicitly when embedding search queries
|
|
41
|
+
document_instruction: str | None = None
|
|
42
|
+
query_instruction: str | None = None
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
EMBEDDING_PARAMS: dict[str, EmbeddingParams] = {
|
|
46
|
+
EmbeddingModel.QWEN_3_VL_8B: EmbeddingParams(
|
|
47
|
+
template=QWEN_3_VL_EMBEDDING_TEMPLATE,
|
|
48
|
+
dimensions=4096,
|
|
49
|
+
document_instruction=EMBEDDING_DOCUMENT_INSTRUCTION,
|
|
50
|
+
query_instruction=EMBEDDING_QUERY_INSTRUCTION,
|
|
51
|
+
),
|
|
52
|
+
}
|
sthai/models.py
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
from sthai.typing import StrEnum
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
class InferenceModel(StrEnum):
|
|
5
|
+
QWEN_3_6_27B = "Qwen/Qwen3.6-27B"
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class EmbeddingModel(StrEnum):
|
|
9
|
+
QWEN_3_VL_8B = "Qwen/Qwen3-VL-Embedding-8B"
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class RerankingModel(StrEnum):
|
|
13
|
+
QWEN_3_VL_8B = "Qwen/Qwen3-VL-Reranker-8B"
|
sthai/py.typed
ADDED
|
File without changes
|