mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 09:25:14 +03:00
feat(typesafe): make classifier questions invocation-scoped (#40659)
makes TypeSafe questions invocation-scoped so each classification request carries both the state being evaluated and the typed questions to answer. We're adding two overloads since we need to comply with Runnable inheritance. * also adds ClassifierRequest + and renames ClassifierResponse to be the public request and response types * normalizes both invocation forms into a complete request before callbacks and tracing begin. * update request serialization, public exports, documentation, and examples for the new request-scoped API. ### Middleware * derive Auto Mode and model-routing questions from validated middleware configuration at invocation time instead of storing mutable question mappings on middleware instances. * pass state and questions explicitly through the classifier keyword API
This commit is contained in:
1 parent
f72f934cef
commit
115dbbd158
12 files changed
+385
-251
No files matched your search
@@ -18,30 +18,46 @@ Set the `TYPESAFE_API_KEY` environment variable before making requests.
|
||||
```python
|
||||
from langchain_typesafe import Choice, Noul, Score, TypeSafeClassifier
|
||||
|
||||
classifier = TypeSafeClassifier(
|
||||
questions={
|
||||
"department": Choice(
|
||||
instructions="Which team should handle this?",
|
||||
criteria={
|
||||
"billing": "Payment or subscription issues",
|
||||
"technical": "Product or integration issues",
|
||||
},
|
||||
),
|
||||
"urgent": Noul(instructions="Does this message express urgency?"),
|
||||
"frustration": Score(
|
||||
instructions="How frustrated does the customer appear?",
|
||||
criteria=["calm", "frustrated", "angry"],
|
||||
),
|
||||
classifier = TypeSafeClassifier()
|
||||
|
||||
result = classifier.invoke(
|
||||
{
|
||||
"state": "Stripe has failed to connect for three days. Help ASAP.",
|
||||
"questions": {
|
||||
"department": Choice(
|
||||
instructions="Which team should handle this?",
|
||||
criteria={
|
||||
"billing": "Payment or subscription issues",
|
||||
"technical": "Product or integration issues",
|
||||
},
|
||||
),
|
||||
"urgent": Noul(instructions="Does this message express urgency?"),
|
||||
"frustration": Score(
|
||||
instructions="How frustrated does the customer appear?",
|
||||
criteria=["calm", "frustrated", "angry"],
|
||||
),
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
result = classifier.invoke("Stripe has failed to connect for three days. Help ASAP.")
|
||||
print(result.choices["department"].choice)
|
||||
print(result.nouls["urgent"].noul)
|
||||
print(result.scores["frustration"].score)
|
||||
```
|
||||
|
||||
Use `await classifier.ainvoke(...)` for asynchronous applications. As a `Runnable`, the classifier can also be composed with other LangChain runnables and supports standard batching, callbacks, and tracing.
|
||||
Pass a complete `ClassifierRequest` mapping to `invoke` or `ainvoke`. Keeping both
|
||||
`state` and `questions` in the Runnable input makes the complete classification request
|
||||
available to composition, batching, callbacks, and tracing. Use
|
||||
`await classifier.ainvoke(...)` for asynchronous applications:
|
||||
|
||||
```python
|
||||
from langchain_typesafe import ClassifierRequest
|
||||
|
||||
request: ClassifierRequest = {
|
||||
"state": "Stripe has failed to connect for three days. Help ASAP.",
|
||||
"questions": {"urgent": Noul(instructions="Is this urgent?")},
|
||||
}
|
||||
result = classifier.invoke(request)
|
||||
```
|
||||
|
||||
### Experimental middleware
|
||||
|
||||
@@ -114,11 +130,16 @@ from langchain_core.messages import HumanMessage, SystemMessage
|
||||
|
||||
response = classifier.invoke(
|
||||
{
|
||||
"conversation": [
|
||||
SystemMessage("You are reviewing a customer support conversation."),
|
||||
HumanMessage("My payouts have failed for three days. Help!"),
|
||||
],
|
||||
"account_tier": "enterprise",
|
||||
"state": {
|
||||
"conversation": [
|
||||
SystemMessage("You are reviewing a customer support conversation."),
|
||||
HumanMessage("My payouts have failed for three days. Help!"),
|
||||
],
|
||||
"account_tier": "enterprise",
|
||||
},
|
||||
"questions": {
|
||||
"urgent": Noul(instructions="Does this customer need urgent help?")
|
||||
},
|
||||
}
|
||||
)
|
||||
```
|
||||
@@ -131,7 +152,6 @@ The classifier creates sync and async `httpx2` clients when they are not supplie
|
||||
import httpx2
|
||||
|
||||
classifier = TypeSafeClassifier(
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
client=httpx2.Client(proxy="http://proxy.internal"),
|
||||
async_client=httpx2.AsyncClient(proxy="http://proxy.internal"),
|
||||
)
|
||||
@@ -148,7 +168,12 @@ from langchain_core.exceptions import ModelAuthenticationError, ModelRateLimitEr
|
||||
from langchain_typesafe import TypeSafeRateLimitError
|
||||
|
||||
try:
|
||||
response = classifier.invoke("Classify this message.")
|
||||
response = classifier.invoke(
|
||||
{
|
||||
"state": "Classify this message.",
|
||||
"questions": {"urgent": Noul(instructions="Is this urgent?")},
|
||||
}
|
||||
)
|
||||
except TypeSafeRateLimitError as error:
|
||||
print(error.request_id, error.retry_after_ms)
|
||||
except (ModelAuthenticationError, ModelRateLimitError):
|
||||
|
||||
@@ -6,6 +6,8 @@ from langchain_typesafe.types import (
|
||||
Answer,
|
||||
Choice,
|
||||
ChoiceAnswer,
|
||||
ClassifierRequest,
|
||||
ClassifierResponse,
|
||||
Noul,
|
||||
NoulAnswer,
|
||||
NoulCriteria,
|
||||
@@ -20,6 +22,8 @@ __all__ = [
|
||||
"Answer",
|
||||
"Choice",
|
||||
"ChoiceAnswer",
|
||||
"ClassifierRequest",
|
||||
"ClassifierResponse",
|
||||
"Noul",
|
||||
"NoulAnswer",
|
||||
"NoulCriteria",
|
||||
|
||||
@@ -28,7 +28,10 @@ from langchain_typesafe.client import (
|
||||
TypeSafeAPITimeoutError,
|
||||
parse_response,
|
||||
)
|
||||
from langchain_typesafe.types import ClassificationResponse, Question, State
|
||||
from langchain_typesafe.types import (
|
||||
ClassifierRequest,
|
||||
ClassifierResponse,
|
||||
)
|
||||
|
||||
_DEFAULT_BASE_URL = "https://api.typesafe.ai"
|
||||
_DEFAULT_MODEL = "jev-latest"
|
||||
@@ -39,7 +42,7 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@beta()
|
||||
class TypeSafeClassifier(RunnableSerializable[State, ClassificationResponse]):
|
||||
class TypeSafeClassifier(RunnableSerializable[ClassifierRequest, ClassifierResponse]):
|
||||
"""Classify JSON-compatible state with TypeSafe.
|
||||
|
||||
`TypeSafeClassifier` is a LangChain `Runnable` for asking one or more typed
|
||||
@@ -69,9 +72,7 @@ class TypeSafeClassifier(RunnableSerializable[State, ClassificationResponse]):
|
||||
constructor values take precedence over environment configuration.
|
||||
|
||||
Args:
|
||||
questions: Named `Noul`, `Choice`, or `Score` questions. Names become keys in
|
||||
`ClassificationResponse.answers`.
|
||||
model: TypeSafe model used to answer the questions.
|
||||
model: TypeSafe model used to answer invocation questions.
|
||||
api_key: TypeSafe API key. If omitted, reads `TYPESAFE_API_KEY`.
|
||||
base_url: Root URL for the TypeSafe API.
|
||||
timeout: Timeout, in seconds, applied to clients created by this class.
|
||||
@@ -79,8 +80,7 @@ class TypeSafeClassifier(RunnableSerializable[State, ClassificationResponse]):
|
||||
async_client: Optional asynchronous `httpx2.AsyncClient` used by `ainvoke`.
|
||||
|
||||
Raises:
|
||||
ValueError: If questions are empty, credentials are unavailable, or the timeout
|
||||
is not positive.
|
||||
ValueError: If credentials are unavailable or the timeout is not positive.
|
||||
|
||||
??? example "Classify state on several dimensions"
|
||||
|
||||
@@ -90,27 +90,31 @@ class TypeSafeClassifier(RunnableSerializable[State, ClassificationResponse]):
|
||||
```python
|
||||
from langchain_typesafe import Choice, Noul, Score, TypeSafeClassifier
|
||||
|
||||
classifier = TypeSafeClassifier(
|
||||
questions={
|
||||
"department": Choice(
|
||||
instructions="Which team should handle this request?",
|
||||
criteria={
|
||||
"billing": "Payment or subscription issues.",
|
||||
"technical": "Product bugs or integration failures.",
|
||||
},
|
||||
),
|
||||
"urgent": Noul(
|
||||
instructions="Does this message require an urgent response?"
|
||||
),
|
||||
"frustration": Score(
|
||||
instructions="How frustrated does the customer appear?",
|
||||
criteria=["Calm.", "Concerned but civil.", "Very angry."],
|
||||
),
|
||||
}
|
||||
)
|
||||
classifier = TypeSafeClassifier()
|
||||
|
||||
response = classifier.invoke(
|
||||
"Stripe has failed to connect for three days. Please help immediately."
|
||||
{
|
||||
"state": (
|
||||
"Stripe has failed to connect for three days. "
|
||||
"Please help immediately."
|
||||
),
|
||||
"questions": {
|
||||
"department": Choice(
|
||||
instructions="Which team should handle this request?",
|
||||
criteria={
|
||||
"billing": "Payment or subscription issues.",
|
||||
"technical": "Product bugs or integration failures.",
|
||||
},
|
||||
),
|
||||
"urgent": Noul(
|
||||
instructions="Does this message require an urgent response?"
|
||||
),
|
||||
"frustration": Score(
|
||||
instructions="How frustrated does the customer appear?",
|
||||
criteria=["Calm.", "Concerned but civil.", "Very angry."],
|
||||
),
|
||||
},
|
||||
}
|
||||
)
|
||||
print(response.choices["department"].choice)
|
||||
print(response.nouls["urgent"].noul)
|
||||
@@ -120,37 +124,27 @@ class TypeSafeClassifier(RunnableSerializable[State, ClassificationResponse]):
|
||||
??? example "Classify asynchronously"
|
||||
|
||||
`ainvoke` uses the classifier's asynchronous HTTP client and returns the same
|
||||
`ClassificationResponse` type as `invoke`.
|
||||
`ClassifierResponse` type as `invoke`.
|
||||
|
||||
```python
|
||||
from langchain_typesafe import Noul, TypeSafeClassifier
|
||||
|
||||
classifier = TypeSafeClassifier(
|
||||
questions={
|
||||
"refund_requested": Noul(
|
||||
instructions="Does the customer request a refund?"
|
||||
)
|
||||
classifier = TypeSafeClassifier()
|
||||
|
||||
response = await classifier.ainvoke(
|
||||
{
|
||||
"state": "Please refund the duplicate charge.",
|
||||
"questions": {
|
||||
"refund_requested": Noul(
|
||||
instructions="Does the customer request a refund?"
|
||||
)
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
response = await classifier.ainvoke("Please refund the duplicate charge.")
|
||||
print(response.nouls["refund_requested"].noul)
|
||||
```
|
||||
"""
|
||||
|
||||
questions: dict[str, Question] = Field(min_length=1)
|
||||
"""Questions sent together for every classifier invocation.
|
||||
|
||||
The mapping key is the question ID and becomes the corresponding key in
|
||||
`ClassificationResponse.answers`. Question IDs identify answers for application
|
||||
code; put the complete judgment in each question's `instructions` rather than
|
||||
relying on its ID to provide model context.
|
||||
|
||||
Questions share the same input state but are evaluated independently. Mix `Choice`,
|
||||
`Noul`, and `Score` questions in one mapping when several judgments use the same
|
||||
state instead of issuing one request per question.
|
||||
"""
|
||||
|
||||
model: str = Field(default=_DEFAULT_MODEL, min_length=1)
|
||||
"""TypeSafe model name used for classification.
|
||||
|
||||
@@ -181,18 +175,13 @@ class TypeSafeClassifier(RunnableSerializable[State, ClassificationResponse]):
|
||||
```python
|
||||
from langchain_typesafe import Noul, TypeSafeClassifier
|
||||
|
||||
classifier = TypeSafeClassifier(
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")}
|
||||
)
|
||||
classifier = TypeSafeClassifier()
|
||||
```
|
||||
|
||||
??? example "Specify directly"
|
||||
|
||||
```python
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key="...",
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
)
|
||||
classifier = TypeSafeClassifier(api_key="...")
|
||||
```
|
||||
"""
|
||||
|
||||
@@ -302,14 +291,14 @@ class TypeSafeClassifier(RunnableSerializable[State, ClassificationResponse]):
|
||||
@override
|
||||
def invoke(
|
||||
self,
|
||||
input: State,
|
||||
input: ClassifierRequest,
|
||||
config: RunnableConfig | None = None,
|
||||
**_: Any,
|
||||
) -> ClassificationResponse:
|
||||
"""Classify one JSON-compatible input synchronously.
|
||||
) -> ClassifierResponse:
|
||||
"""Classify one request synchronously.
|
||||
|
||||
Args:
|
||||
input: Text, object, array, `BaseMessage`, or message sequence to classify.
|
||||
input: Complete request containing the state and typed questions.
|
||||
config: Optional LangChain runnable configuration for callbacks, tags,
|
||||
metadata, and tracing.
|
||||
**_: Additional keyword arguments accepted for `Runnable` compatibility and
|
||||
@@ -319,7 +308,6 @@ class TypeSafeClassifier(RunnableSerializable[State, ClassificationResponse]):
|
||||
Structured TypeSafe answers and request metadata.
|
||||
|
||||
Raises:
|
||||
TypeError: If the input is not a supported state.
|
||||
TypeSafeAPIError: If TypeSafe returns an unsuccessful HTTP response.
|
||||
TypeSafeAPIConnectionError: If no HTTP response is received.
|
||||
TypeSafeAPITimeoutError: If the request exceeds its client timeout.
|
||||
@@ -335,14 +323,14 @@ class TypeSafeClassifier(RunnableSerializable[State, ClassificationResponse]):
|
||||
@override
|
||||
async def ainvoke(
|
||||
self,
|
||||
input: State,
|
||||
input: ClassifierRequest,
|
||||
config: RunnableConfig | None = None,
|
||||
**_: Any,
|
||||
) -> ClassificationResponse:
|
||||
"""Classify one JSON-compatible input asynchronously.
|
||||
) -> ClassifierResponse:
|
||||
"""Classify one request asynchronously.
|
||||
|
||||
Args:
|
||||
input: Text, object, array, `BaseMessage`, or message sequence to classify.
|
||||
input: Complete request containing the state and typed questions.
|
||||
config: Optional LangChain runnable configuration for callbacks, tags,
|
||||
metadata, and tracing.
|
||||
**_: Additional keyword arguments accepted for `Runnable` compatibility and
|
||||
@@ -352,7 +340,6 @@ class TypeSafeClassifier(RunnableSerializable[State, ClassificationResponse]):
|
||||
Structured TypeSafe answers and request metadata.
|
||||
|
||||
Raises:
|
||||
TypeError: If the input is not a supported state.
|
||||
TypeSafeAPIError: If TypeSafe returns an unsuccessful HTTP response.
|
||||
TypeSafeAPIConnectionError: If no HTTP response is received.
|
||||
TypeSafeAPITimeoutError: If the request exceeds its client timeout.
|
||||
@@ -365,8 +352,8 @@ class TypeSafeClassifier(RunnableSerializable[State, ClassificationResponse]):
|
||||
run_type="llm",
|
||||
)
|
||||
|
||||
def _classify(self, state: State) -> ClassificationResponse:
|
||||
payload = self._payload(state)
|
||||
def _classify(self, request: ClassifierRequest) -> ClassifierResponse:
|
||||
payload = self._payload(request)
|
||||
if self.client is None: # pragma: no cover - guaranteed by model validation
|
||||
message = "Synchronous TypeSafe client was not initialized."
|
||||
raise TypeSafeAPIConnectionError(message)
|
||||
@@ -383,8 +370,11 @@ class TypeSafeClassifier(RunnableSerializable[State, ClassificationResponse]):
|
||||
raise TypeSafeAPIConnectionError(message) from error
|
||||
return self._record_usage(parse_response(response))
|
||||
|
||||
async def _aclassify(self, state: State) -> ClassificationResponse:
|
||||
payload = self._payload(state)
|
||||
async def _aclassify(
|
||||
self,
|
||||
request: ClassifierRequest,
|
||||
) -> ClassifierResponse:
|
||||
payload = self._payload(request)
|
||||
if self.async_client is None: # pragma: no cover - guaranteed by validation
|
||||
message = "Asynchronous TypeSafe client was not initialized."
|
||||
raise TypeSafeAPIConnectionError(message)
|
||||
@@ -412,7 +402,7 @@ class TypeSafeClassifier(RunnableSerializable[State, ClassificationResponse]):
|
||||
}
|
||||
return config
|
||||
|
||||
def _record_usage(self, response: ClassificationResponse) -> ClassificationResponse:
|
||||
def _record_usage(self, response: ClassifierResponse) -> ClassifierResponse:
|
||||
"""Attach TypeSafe token usage to the active run, if there is one.
|
||||
|
||||
Nothing is written when tracing is disabled, and a tracing failure never fails
|
||||
@@ -449,13 +439,13 @@ class TypeSafeClassifier(RunnableSerializable[State, ClassificationResponse]):
|
||||
"User-Agent": f"langchain-typesafe/{__version__}",
|
||||
}
|
||||
|
||||
def _payload(self, state: State) -> dict[str, JsonValue]:
|
||||
def _payload(self, request: ClassifierRequest) -> dict[str, JsonValue]:
|
||||
return {
|
||||
"state": serialize_state(state),
|
||||
"state": serialize_state(request["state"]),
|
||||
"model": self.model,
|
||||
"questions": {
|
||||
name: question.model_dump(mode="json", exclude_none=True)
|
||||
for name, question in self.questions.items()
|
||||
for name, question in request["questions"].items()
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -23,7 +23,7 @@ from langchain_core.exceptions import (
|
||||
from pydantic import ValidationError
|
||||
from typing_extensions import override
|
||||
|
||||
from langchain_typesafe.types import ClassificationResponse
|
||||
from langchain_typesafe.types import ClassifierResponse
|
||||
|
||||
_REQUEST_ID_HEADER = "x-typesafe-request-id"
|
||||
_RETRY_AFTER_HEADER = "retry-after"
|
||||
@@ -310,7 +310,7 @@ def _api_error(response: httpx2.Response) -> TypeSafeAPIError:
|
||||
)
|
||||
|
||||
|
||||
def parse_response(response: httpx2.Response) -> ClassificationResponse:
|
||||
def parse_response(response: httpx2.Response) -> ClassifierResponse:
|
||||
"""Validate an HTTP response and convert it to a classification response.
|
||||
|
||||
Args:
|
||||
@@ -331,7 +331,7 @@ def parse_response(response: httpx2.Response) -> ClassificationResponse:
|
||||
endpoint = _response_endpoint(response)
|
||||
body = _response_body(response)
|
||||
try:
|
||||
parsed = ClassificationResponse.model_validate(body)
|
||||
parsed = ClassifierResponse.model_validate(body)
|
||||
except ValidationError as error:
|
||||
location = error.errors()[0].get("loc", ())
|
||||
field_path = ".".join(str(item) for item in location) or "response"
|
||||
|
||||
@@ -28,7 +28,7 @@ from pydantic import BaseModel, Field
|
||||
from typing_extensions import override
|
||||
|
||||
from langchain_typesafe.classifier import TypeSafeClassifier
|
||||
from langchain_typesafe.types import Noul, NoulCriteria
|
||||
from langchain_typesafe.types import Noul, NoulCriteria, Question
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Awaitable, Callable
|
||||
@@ -71,6 +71,16 @@ class _AutoModeConfig(BaseModel):
|
||||
)
|
||||
|
||||
|
||||
def _risk_questions(config: _AutoModeConfig) -> dict[str, Question]:
|
||||
"""Build the risk question from validated middleware configuration."""
|
||||
return {
|
||||
_QUESTION_ID: Noul(
|
||||
instructions=config.instructions,
|
||||
criteria=config.criteria,
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
class AutoModeMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, ResponseT]):
|
||||
"""Allow low-risk tool calls and block risky calls using TypeSafe.
|
||||
|
||||
@@ -154,14 +164,7 @@ class AutoModeMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, Respon
|
||||
"criteria": criteria,
|
||||
}
|
||||
)
|
||||
self.classifier = TypeSafeClassifier(
|
||||
questions={
|
||||
_QUESTION_ID: Noul(
|
||||
instructions=self.config.instructions,
|
||||
criteria=self.config.criteria,
|
||||
)
|
||||
},
|
||||
)
|
||||
self.classifier = TypeSafeClassifier()
|
||||
|
||||
@staticmethod
|
||||
def _classification_state(request: ToolCallRequest) -> dict[str, Any]:
|
||||
@@ -219,7 +222,12 @@ class AutoModeMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, Respon
|
||||
"""
|
||||
if request.tool_call["name"] not in self._tool_names:
|
||||
return handler(request)
|
||||
response = self.classifier.invoke(self._classification_state(request))
|
||||
response = self.classifier.invoke(
|
||||
{
|
||||
"state": self._classification_state(request),
|
||||
"questions": _risk_questions(self.config),
|
||||
}
|
||||
)
|
||||
probability = response.nouls[_QUESTION_ID].noul
|
||||
if probability >= _PROBABILITY_THRESHOLD:
|
||||
return self._blocked_tool_message(request, probability)
|
||||
@@ -245,7 +253,12 @@ class AutoModeMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, Respon
|
||||
"""
|
||||
if request.tool_call["name"] not in self._tool_names:
|
||||
return await handler(request)
|
||||
response = await self.classifier.ainvoke(self._classification_state(request))
|
||||
response = await self.classifier.ainvoke(
|
||||
{
|
||||
"state": self._classification_state(request),
|
||||
"questions": _risk_questions(self.config),
|
||||
}
|
||||
)
|
||||
probability = response.nouls[_QUESTION_ID].noul
|
||||
if probability >= _PROBABILITY_THRESHOLD:
|
||||
return self._blocked_tool_message(request, probability)
|
||||
|
||||
@@ -32,7 +32,7 @@ from pydantic import BaseModel, Field, JsonValue
|
||||
from typing_extensions import NotRequired, override
|
||||
|
||||
from langchain_typesafe.classifier import TypeSafeClassifier
|
||||
from langchain_typesafe.types import Choice, ChoiceAnswer
|
||||
from langchain_typesafe.types import Choice, ChoiceAnswer, Question
|
||||
|
||||
_QUESTION_ID = "model_route"
|
||||
_QuestionContent = str | dict[str, JsonValue] | list[JsonValue]
|
||||
@@ -58,6 +58,18 @@ class _ModelRouterConfig(BaseModel):
|
||||
instructions: _QuestionContent
|
||||
|
||||
|
||||
def _routing_questions(config: _ModelRouterConfig) -> dict[str, Question]:
|
||||
"""Build the routing question from validated middleware configuration."""
|
||||
return {
|
||||
_QUESTION_ID: Choice(
|
||||
instructions=config.instructions,
|
||||
criteria={
|
||||
route: choice.criteria for route, choice in config.choices.items()
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
class _ModelRouterState(AgentState):
|
||||
"""Agent state used to persist the TypeSafe routing answer."""
|
||||
|
||||
@@ -137,17 +149,7 @@ class ModelRouterMiddleware(AgentMiddleware[_ModelRouterState]):
|
||||
else choice.model
|
||||
for route, choice in self.config.choices.items()
|
||||
}
|
||||
self.classifier = TypeSafeClassifier(
|
||||
questions={
|
||||
_QUESTION_ID: Choice(
|
||||
instructions=self.config.instructions,
|
||||
criteria={
|
||||
route: choice.criteria
|
||||
for route, choice in self.config.choices.items()
|
||||
},
|
||||
)
|
||||
}
|
||||
)
|
||||
self.classifier = TypeSafeClassifier()
|
||||
|
||||
@staticmethod
|
||||
def _latest_human_message(state: _ModelRouterState) -> HumanMessage:
|
||||
@@ -163,7 +165,12 @@ class ModelRouterMiddleware(AgentMiddleware[_ModelRouterState]):
|
||||
self, state: _ModelRouterState, runtime: Runtime[ContextT]
|
||||
) -> dict[str, ChoiceAnswer]:
|
||||
"""Classify the latest task and store the complete routing answer."""
|
||||
response = self.classifier.invoke(self._latest_human_message(state))
|
||||
response = self.classifier.invoke(
|
||||
{
|
||||
"state": self._latest_human_message(state),
|
||||
"questions": _routing_questions(self.config),
|
||||
}
|
||||
)
|
||||
return {"model_route": response.choices[_QUESTION_ID]}
|
||||
|
||||
@override
|
||||
@@ -171,7 +178,12 @@ class ModelRouterMiddleware(AgentMiddleware[_ModelRouterState]):
|
||||
self, state: _ModelRouterState, runtime: Runtime[ContextT]
|
||||
) -> dict[str, ChoiceAnswer]:
|
||||
"""Classify the latest task asynchronously and store the routing answer."""
|
||||
response = await self.classifier.ainvoke(self._latest_human_message(state))
|
||||
response = await self.classifier.ainvoke(
|
||||
{
|
||||
"state": self._latest_human_message(state),
|
||||
"questions": _routing_questions(self.config),
|
||||
}
|
||||
)
|
||||
return {"model_route": response.choices[_QUESTION_ID]}
|
||||
|
||||
@override
|
||||
|
||||
@@ -7,6 +7,7 @@ from typing import Annotated, Literal, TypeAlias
|
||||
|
||||
from langchain_core.messages import BaseMessage
|
||||
from pydantic import BaseModel, ConfigDict, Field, JsonValue
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
_QuestionContent: TypeAlias = str | dict[str, JsonValue] | list[JsonValue]
|
||||
_StateValue: TypeAlias = (
|
||||
@@ -61,14 +62,17 @@ class Noul(BaseModel):
|
||||
```python
|
||||
from langchain_typesafe import Noul, TypeSafeClassifier
|
||||
|
||||
classifier = TypeSafeClassifier(
|
||||
questions={
|
||||
"urgent": Noul(
|
||||
instructions="Does this message require an urgent response?"
|
||||
)
|
||||
classifier = TypeSafeClassifier()
|
||||
response = classifier.invoke(
|
||||
{
|
||||
"state": "Production is down. Please help immediately.",
|
||||
"questions": {
|
||||
"urgent": Noul(
|
||||
instructions="Does this message require an urgent response?"
|
||||
)
|
||||
},
|
||||
}
|
||||
)
|
||||
response = classifier.invoke("Production is down. Please help immediately.")
|
||||
urgency = response.nouls["urgent"].noul
|
||||
|
||||
if urgency >= 0.8:
|
||||
@@ -83,7 +87,7 @@ class Noul(BaseModel):
|
||||
"""Complete yes/no judgment to make about the input state.
|
||||
|
||||
Instructions may be text or structured JSON. Write the full question here even when
|
||||
the question ID used by `TypeSafeClassifier.questions` appears self-explanatory.
|
||||
its ID in `ClassifierRequest.questions` appears self-explanatory.
|
||||
"""
|
||||
|
||||
criteria: NoulCriteria | None = None
|
||||
@@ -103,19 +107,22 @@ class Choice(BaseModel):
|
||||
```python
|
||||
from langchain_typesafe import Choice, TypeSafeClassifier
|
||||
|
||||
classifier = TypeSafeClassifier(
|
||||
questions={
|
||||
"department": Choice(
|
||||
instructions="Which team should handle this request?",
|
||||
criteria={
|
||||
"billing": "Payment, invoice, or subscription issues.",
|
||||
"technical": "Product bugs or integration failures.",
|
||||
"sales": "Pricing or purchasing questions.",
|
||||
},
|
||||
)
|
||||
classifier = TypeSafeClassifier()
|
||||
response = classifier.invoke(
|
||||
{
|
||||
"state": "Stripe fails whenever I connect my account.",
|
||||
"questions": {
|
||||
"department": Choice(
|
||||
instructions="Which team should handle this request?",
|
||||
criteria={
|
||||
"billing": "Payment, invoice, or subscription issues.",
|
||||
"technical": "Product bugs or integration failures.",
|
||||
"sales": "Pricing or purchasing questions.",
|
||||
},
|
||||
)
|
||||
},
|
||||
}
|
||||
)
|
||||
response = classifier.invoke("Stripe fails whenever I connect my account.")
|
||||
department = response.choices["department"]
|
||||
|
||||
if department.confidence >= 0.7:
|
||||
@@ -153,19 +160,22 @@ class Score(BaseModel):
|
||||
```python
|
||||
from langchain_typesafe import Score, TypeSafeClassifier
|
||||
|
||||
classifier = TypeSafeClassifier(
|
||||
questions={
|
||||
"frustration": Score(
|
||||
instructions="How frustrated does the customer appear?",
|
||||
criteria=[
|
||||
"Calm and neutral.",
|
||||
"Concerned but civil.",
|
||||
"Very angry or using strong language.",
|
||||
],
|
||||
)
|
||||
classifier = TypeSafeClassifier()
|
||||
response = classifier.invoke(
|
||||
{
|
||||
"state": "This has failed three times. Fix it now.",
|
||||
"questions": {
|
||||
"frustration": Score(
|
||||
instructions="How frustrated does the customer appear?",
|
||||
criteria=[
|
||||
"Calm and neutral.",
|
||||
"Concerned but civil.",
|
||||
"Very angry or using strong language.",
|
||||
],
|
||||
)
|
||||
},
|
||||
}
|
||||
)
|
||||
response = classifier.invoke("This has failed three times. Fix it now.")
|
||||
frustration = response.scores["frustration"]
|
||||
|
||||
print(frustration.score) # May be fractional, for example 1.35.
|
||||
@@ -188,6 +198,20 @@ Question = Annotated[Noul | Choice | Score, Field(discriminator="type")]
|
||||
"""A discriminated union of question types accepted by `TypeSafeClassifier`."""
|
||||
|
||||
|
||||
class ClassifierRequest(TypedDict):
|
||||
"""Complete input for one `TypeSafeClassifier` invocation.
|
||||
|
||||
Keeping the state and questions in the Runnable input ensures both values
|
||||
participate in composition, batching, and tracing.
|
||||
"""
|
||||
|
||||
state: State
|
||||
"""Text, structured JSON, or LangChain messages to classify."""
|
||||
|
||||
questions: dict[str, Question]
|
||||
"""Non-empty mapping of answer IDs to typed classification questions."""
|
||||
|
||||
|
||||
class NoulAnswer(BaseModel):
|
||||
"""Probability that a `Noul` question's answer is yes."""
|
||||
|
||||
@@ -260,12 +284,12 @@ class Usage(BaseModel):
|
||||
"""Number of output tokens produced, or `None` when not reported."""
|
||||
|
||||
|
||||
class ClassificationResponse(BaseModel):
|
||||
class ClassifierResponse(BaseModel):
|
||||
"""Typed answers and metadata returned from one TypeSafe request.
|
||||
|
||||
Access every answer through `answers`, or use `nouls`, `choices`, and `scores` for
|
||||
views filtered by answer type. Each mapping preserves the question IDs supplied to
|
||||
`TypeSafeClassifier.questions`.
|
||||
`ClassifierRequest.questions`.
|
||||
"""
|
||||
|
||||
model: str
|
||||
@@ -312,7 +336,8 @@ __all__ = [
|
||||
"Answer",
|
||||
"Choice",
|
||||
"ChoiceAnswer",
|
||||
"ClassificationResponse",
|
||||
"ClassifierRequest",
|
||||
"ClassifierResponse",
|
||||
"Noul",
|
||||
"NoulAnswer",
|
||||
"NoulCriteria",
|
||||
|
||||
@@ -10,8 +10,10 @@ from langchain_core.messages import HumanMessage, SystemMessage
|
||||
from langchain_typesafe import (
|
||||
Choice,
|
||||
ChoiceAnswer,
|
||||
ClassifierRequest,
|
||||
Noul,
|
||||
NoulAnswer,
|
||||
Question,
|
||||
Score,
|
||||
ScoreAnswer,
|
||||
TypeSafeClassifier,
|
||||
@@ -21,40 +23,39 @@ from langchain_typesafe import (
|
||||
def test_invoke_all_question_types() -> None:
|
||||
"""Exercise the live sync API across Choice, Noul, and Score questions."""
|
||||
labels = {"billing", "technical", "sales"}
|
||||
classifier = TypeSafeClassifier(
|
||||
questions={
|
||||
"department": Choice(
|
||||
instructions="Which team should handle this request?",
|
||||
criteria={
|
||||
"billing": "Payment or subscription issues.",
|
||||
"technical": "Product bugs or integration failures.",
|
||||
"sales": "Pricing or purchasing questions.",
|
||||
},
|
||||
),
|
||||
"urgent": Noul(
|
||||
instructions="Does this message require an urgent response?"
|
||||
),
|
||||
"frustration": Score(
|
||||
instructions="How frustrated does the customer appear?",
|
||||
criteria=[
|
||||
"Calm and neutral.",
|
||||
"Concerned but civil.",
|
||||
"Very angry or using strong language.",
|
||||
],
|
||||
),
|
||||
}
|
||||
)
|
||||
questions: dict[str, Question] = {
|
||||
"department": Choice(
|
||||
instructions="Which team should handle this request?",
|
||||
criteria={
|
||||
"billing": "Payment or subscription issues.",
|
||||
"technical": "Product bugs or integration failures.",
|
||||
"sales": "Pricing or purchasing questions.",
|
||||
},
|
||||
),
|
||||
"urgent": Noul(instructions="Does this message require an urgent response?"),
|
||||
"frustration": Score(
|
||||
instructions="How frustrated does the customer appear?",
|
||||
criteria=[
|
||||
"Calm and neutral.",
|
||||
"Concerned but civil.",
|
||||
"Very angry or using strong language.",
|
||||
],
|
||||
),
|
||||
}
|
||||
classifier = TypeSafeClassifier()
|
||||
|
||||
try:
|
||||
response = classifier.invoke(
|
||||
{
|
||||
request: ClassifierRequest = {
|
||||
"state": {
|
||||
"message": (
|
||||
"Stripe has failed to connect for three days. "
|
||||
"Please help immediately."
|
||||
),
|
||||
"account_tier": "enterprise",
|
||||
}
|
||||
)
|
||||
},
|
||||
"questions": questions,
|
||||
}
|
||||
response = classifier.invoke(request)
|
||||
|
||||
department = response.answers["department"]
|
||||
urgent = response.answers["urgent"]
|
||||
@@ -87,17 +88,16 @@ def test_invoke_all_question_types() -> None:
|
||||
|
||||
async def test_ainvoke_with_nested_messages() -> None:
|
||||
"""Exercise the live async API with messages nested in structured state."""
|
||||
classifier = TypeSafeClassifier(
|
||||
questions={
|
||||
"needs_support": Noul(
|
||||
instructions="Does the user need help resolving a technical problem?"
|
||||
)
|
||||
}
|
||||
)
|
||||
questions: dict[str, Question] = {
|
||||
"needs_support": Noul(
|
||||
instructions="Does the user need help resolving a technical problem?"
|
||||
)
|
||||
}
|
||||
classifier = TypeSafeClassifier()
|
||||
|
||||
try:
|
||||
response = await classifier.ainvoke(
|
||||
{
|
||||
request: ClassifierRequest = {
|
||||
"state": {
|
||||
"conversation": [
|
||||
SystemMessage("You are reviewing a customer support conversation."),
|
||||
HumanMessage(
|
||||
@@ -111,8 +111,10 @@ async def test_ainvoke_with_nested_messages() -> None:
|
||||
"trial": False,
|
||||
"notes": None,
|
||||
},
|
||||
}
|
||||
)
|
||||
},
|
||||
"questions": questions,
|
||||
}
|
||||
response = await classifier.ainvoke(request)
|
||||
|
||||
needs_support = response.answers["needs_support"]
|
||||
assert isinstance(needs_support, NoulAnswer)
|
||||
|
||||
@@ -24,6 +24,7 @@ from langchain_typesafe import NoulCriteria, experimental
|
||||
from langchain_typesafe.client import TypeSafeInternalServerError
|
||||
from langchain_typesafe.experimental.middleware import AutoModeMiddleware
|
||||
from langchain_typesafe.experimental.middleware import __all__ as middleware_all
|
||||
from langchain_typesafe.experimental.middleware.auto_mode import _risk_questions
|
||||
from langchain_typesafe.types import Noul
|
||||
|
||||
API_KEY = "test-api-key"
|
||||
@@ -174,7 +175,7 @@ async def test_middleware_constructs_configurable_risk_classifier() -> None:
|
||||
instructions="Assess production impact.",
|
||||
criteria=custom_criteria,
|
||||
) as middleware:
|
||||
question = middleware.classifier.questions["is_risky"]
|
||||
question = _risk_questions(middleware.config)["is_risky"]
|
||||
|
||||
assert question == Noul(
|
||||
instructions="Assess production impact.",
|
||||
@@ -186,7 +187,7 @@ async def test_none_criteria_is_supported() -> None:
|
||||
"""Allow callers to classify without outcome criteria."""
|
||||
async with _middleware(0.2, tools=["delete_file"]) as middleware:
|
||||
assert middleware.config.criteria is None
|
||||
assert middleware.classifier.questions["is_risky"].criteria is None
|
||||
assert _risk_questions(middleware.config)["is_risky"].criteria is None
|
||||
|
||||
|
||||
async def test_base_tool_name_is_inferred() -> None:
|
||||
|
||||
+18
-8
@@ -18,11 +18,12 @@ from langchain_typesafe.experimental.middleware import (
|
||||
from langchain_typesafe.experimental.middleware import (
|
||||
__all__ as middleware_all,
|
||||
)
|
||||
from langchain_typesafe.types import ClassificationResponse
|
||||
from langchain_typesafe.experimental.middleware.model_router import _routing_questions
|
||||
from langchain_typesafe.types import ClassifierResponse
|
||||
|
||||
|
||||
def _response(route: str) -> ClassificationResponse:
|
||||
return ClassificationResponse(
|
||||
def _response(route: str) -> ClassifierResponse:
|
||||
return ClassifierResponse(
|
||||
model="jev-latest",
|
||||
answers={
|
||||
"model_route": ChoiceAnswer(
|
||||
@@ -72,9 +73,8 @@ def test_middleware_constructs_classifier_from_routing_configuration() -> None:
|
||||
"""Construct a TypeSafe Choice and expose validated configuration fields."""
|
||||
middleware, _, classifier, classifier_class = _router()
|
||||
|
||||
classifier_class.assert_called_once()
|
||||
questions = classifier_class.call_args.kwargs["questions"]
|
||||
assert questions == {
|
||||
classifier_class.assert_called_once_with()
|
||||
assert _routing_questions(middleware.config) == {
|
||||
"model_route": Choice(
|
||||
instructions="Choose the least costly model suited to the task.",
|
||||
criteria={"fast": "Simple tasks.", "powerful": "Complex tasks."},
|
||||
@@ -104,10 +104,20 @@ async def test_agent_routes_using_latest_human_message(*, asynchronous: bool) ->
|
||||
|
||||
if asynchronous:
|
||||
result = await agent.ainvoke(inputs)
|
||||
classifier.ainvoke.assert_awaited_once_with(latest_message)
|
||||
classifier.ainvoke.assert_awaited_once_with(
|
||||
{
|
||||
"state": latest_message,
|
||||
"questions": _routing_questions(middleware.config),
|
||||
}
|
||||
)
|
||||
else:
|
||||
result = agent.invoke(inputs)
|
||||
classifier.invoke.assert_called_once_with(latest_message)
|
||||
classifier.invoke.assert_called_once_with(
|
||||
{
|
||||
"state": latest_message,
|
||||
"questions": _routing_questions(middleware.config),
|
||||
}
|
||||
)
|
||||
|
||||
assert result["messages"][-1].text == "fast response"
|
||||
assert result["model_route"] == _response("fast").choices["model_route"]
|
||||
|
||||
@@ -15,6 +15,7 @@ from pydantic import SecretStr, ValidationError
|
||||
from langchain_typesafe import (
|
||||
Choice,
|
||||
ChoiceAnswer,
|
||||
ClassifierRequest,
|
||||
Noul,
|
||||
NoulAnswer,
|
||||
Score,
|
||||
@@ -43,9 +44,11 @@ class _RunRecorder(BaseCallbackHandler):
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.metadata: dict[str, Any] = {}
|
||||
self.input: Any = None
|
||||
self.run_type: str | None = None
|
||||
|
||||
def on_chain_start(self, *_: Any, **kwargs: Any) -> None:
|
||||
def on_chain_start(self, *args: Any, **kwargs: Any) -> None:
|
||||
self.input = args[1]
|
||||
self.metadata = kwargs.get("metadata") or {}
|
||||
self.run_type = kwargs.get("run_type")
|
||||
|
||||
@@ -91,6 +94,10 @@ def _questions() -> dict[str, Choice | Noul | Score]:
|
||||
}
|
||||
|
||||
|
||||
def _request(state: Any = "hello") -> ClassifierRequest:
|
||||
return {"state": state, "questions": _questions()}
|
||||
|
||||
|
||||
def test_classifier_is_beta() -> None:
|
||||
"""Constructing the classifier warns that its API is in beta."""
|
||||
with pytest.warns(
|
||||
@@ -99,7 +106,6 @@ def test_classifier_is_beta() -> None:
|
||||
):
|
||||
TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
)
|
||||
|
||||
|
||||
@@ -126,7 +132,6 @@ def test_model_must_not_be_empty(model: str) -> None:
|
||||
TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
model=model,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
)
|
||||
|
||||
|
||||
@@ -170,11 +175,14 @@ def test_invoke_sends_request_and_parses_response() -> None:
|
||||
client = httpx2.Client(transport=httpx2.MockTransport(handler))
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions=_questions(),
|
||||
client=client,
|
||||
)
|
||||
|
||||
result = classifier.invoke({"message": "Stripe fails to connect."})
|
||||
request: ClassifierRequest = {
|
||||
"state": {"message": "Stripe fails to connect."},
|
||||
"questions": _questions(),
|
||||
}
|
||||
result = classifier.invoke(request)
|
||||
|
||||
assert result.request_id == REQUEST_ID
|
||||
assert result.usage.input_tokens == 42
|
||||
@@ -207,11 +215,10 @@ def test_single_message_is_serialized_as_role_content_state() -> None:
|
||||
client = httpx2.Client(transport=httpx2.MockTransport(handler))
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
client=client,
|
||||
)
|
||||
|
||||
classifier.invoke(HumanMessage("Please help immediately."))
|
||||
classifier.invoke(_request(HumanMessage("Please help immediately.")))
|
||||
|
||||
assert observed_state == {
|
||||
"role": "user",
|
||||
@@ -232,16 +239,17 @@ def test_message_sequence_is_serialized_as_conversation_state() -> None:
|
||||
client = httpx2.Client(transport=httpx2.MockTransport(handler))
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions={"urgent": Noul(instructions="Is the user asking for help?")},
|
||||
client=client,
|
||||
)
|
||||
|
||||
classifier.invoke(
|
||||
[
|
||||
SystemMessage("You are a support assistant."),
|
||||
HumanMessage("My integration is broken."),
|
||||
AIMessage("I can help troubleshoot it."),
|
||||
]
|
||||
_request(
|
||||
[
|
||||
SystemMessage("You are a support assistant."),
|
||||
HumanMessage("My integration is broken."),
|
||||
AIMessage("I can help troubleshoot it."),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
assert observed_state == [
|
||||
@@ -252,6 +260,59 @@ def test_message_sequence_is_serialized_as_conversation_state() -> None:
|
||||
client.close()
|
||||
|
||||
|
||||
def test_invoke_accepts_classifier_request() -> None:
|
||||
"""The complete typed request is accepted as the Runnable input."""
|
||||
observed_payload: dict[str, Any] = {}
|
||||
questions: dict[str, Choice | Noul | Score] = {
|
||||
"urgent": Noul(instructions="Is this urgent?")
|
||||
}
|
||||
|
||||
def handler(request: httpx2.Request) -> httpx2.Response:
|
||||
observed_payload.update(json.loads(request.content))
|
||||
return httpx2.Response(200, json=_response_payload())
|
||||
|
||||
client = httpx2.Client(transport=httpx2.MockTransport(handler))
|
||||
classifier = TypeSafeClassifier(api_key=API_KEY, client=client)
|
||||
|
||||
request: ClassifierRequest = {
|
||||
"state": {"message": "Please help ASAP."},
|
||||
"questions": questions,
|
||||
}
|
||||
classifier.invoke(request)
|
||||
|
||||
assert observed_payload["state"] == {"message": "Please help ASAP."}
|
||||
assert observed_payload["questions"] == {
|
||||
"urgent": {"type": "noul", "instructions": "Is this urgent?"}
|
||||
}
|
||||
client.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ainvoke_accepts_classifier_request() -> None:
|
||||
"""The asynchronous API accepts the same typed request input."""
|
||||
observed_payload: dict[str, Any] = {}
|
||||
questions: dict[str, Choice | Noul | Score] = {
|
||||
"urgent": Noul(instructions="Is this urgent?")
|
||||
}
|
||||
|
||||
async def handler(request: httpx2.Request) -> httpx2.Response:
|
||||
observed_payload.update(json.loads(request.content))
|
||||
return httpx2.Response(200, json=_response_payload())
|
||||
|
||||
async_client = httpx2.AsyncClient(transport=httpx2.MockTransport(handler))
|
||||
classifier = TypeSafeClassifier(api_key=API_KEY, async_client=async_client)
|
||||
|
||||
request: ClassifierRequest = {
|
||||
"state": "Please help ASAP.",
|
||||
"questions": questions,
|
||||
}
|
||||
await classifier.ainvoke(request)
|
||||
|
||||
assert observed_payload["state"] == "Please help ASAP."
|
||||
assert set(observed_payload["questions"]) == {"urgent"}
|
||||
await async_client.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ainvoke_uses_async_client() -> None:
|
||||
"""The async runnable sends requests through the injected async client."""
|
||||
@@ -263,11 +324,10 @@ async def test_ainvoke_uses_async_client() -> None:
|
||||
async_client = httpx2.AsyncClient(transport=httpx2.MockTransport(handler))
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions=_questions(),
|
||||
async_client=async_client,
|
||||
)
|
||||
|
||||
result = await classifier.ainvoke("Please help ASAP.")
|
||||
result = await classifier.ainvoke(_request("Please help ASAP."))
|
||||
|
||||
assert result.choices["department"].choice == "technical"
|
||||
await async_client.aclose()
|
||||
@@ -295,7 +355,6 @@ async def test_missing_clients_are_created(
|
||||
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
timeout=12.5,
|
||||
)
|
||||
|
||||
@@ -314,7 +373,6 @@ async def test_injected_clients_are_preserved() -> None:
|
||||
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
client=client,
|
||||
async_client=async_client,
|
||||
)
|
||||
@@ -338,12 +396,11 @@ async def test_ainvoke_translates_api_error() -> None:
|
||||
async_client = httpx2.AsyncClient(transport=httpx2.MockTransport(handler))
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
async_client=async_client,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeSafeAPIError) as exc_info:
|
||||
await classifier.ainvoke("hello")
|
||||
await classifier.ainvoke(_request())
|
||||
|
||||
assert exc_info.value.status_code == 429
|
||||
assert exc_info.value.request_id == REQUEST_ID
|
||||
@@ -361,12 +418,11 @@ async def test_ainvoke_translates_connection_error() -> None:
|
||||
async_client = httpx2.AsyncClient(transport=httpx2.MockTransport(handler))
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
async_client=async_client,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeSafeAPIConnectionError, match="Unable to connect"):
|
||||
await classifier.ainvoke("hello")
|
||||
await classifier.ainvoke(_request())
|
||||
|
||||
await async_client.aclose()
|
||||
|
||||
@@ -385,12 +441,11 @@ async def test_ainvoke_translates_timeout_error() -> None:
|
||||
)
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
async_client=async_client,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeSafeAPITimeoutError) as exc_info:
|
||||
await classifier.ainvoke("hello")
|
||||
await classifier.ainvoke(_request())
|
||||
|
||||
assert exc_info.value.timeout == async_client.timeout
|
||||
await async_client.aclose()
|
||||
@@ -399,9 +454,7 @@ async def test_ainvoke_translates_timeout_error() -> None:
|
||||
def test_api_key_from_environment(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""The classifier reads its API key from `TYPESAFE_API_KEY`."""
|
||||
monkeypatch.setenv("TYPESAFE_API_KEY", API_KEY)
|
||||
classifier = TypeSafeClassifier(
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")}
|
||||
)
|
||||
classifier = TypeSafeClassifier()
|
||||
assert isinstance(classifier.api_key, SecretStr)
|
||||
assert classifier.api_key.get_secret_value() == API_KEY
|
||||
|
||||
@@ -412,7 +465,6 @@ def test_base_url_from_environment(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
)
|
||||
|
||||
assert classifier.base_url == "https://gateway.typesafe.example"
|
||||
@@ -427,7 +479,6 @@ def test_explicit_base_url_overrides_environment(
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
base_url="https://explicit.example",
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
)
|
||||
|
||||
assert classifier.base_url == "https://explicit.example"
|
||||
@@ -437,7 +488,7 @@ def test_missing_api_key_is_rejected(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Constructing a classifier without credentials fails before creating clients."""
|
||||
monkeypatch.delenv("TYPESAFE_API_KEY", raising=False)
|
||||
with pytest.raises(ValidationError, match="TypeSafe API key is required"):
|
||||
TypeSafeClassifier(questions={"urgent": Noul(instructions="Is this urgent?")})
|
||||
TypeSafeClassifier()
|
||||
|
||||
|
||||
def test_api_error_does_not_expose_response_body() -> None:
|
||||
@@ -453,12 +504,11 @@ def test_api_error_does_not_expose_response_body() -> None:
|
||||
client = httpx2.Client(transport=httpx2.MockTransport(handler))
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
client=client,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeSafeAPIError) as exc_info:
|
||||
classifier.invoke("hello")
|
||||
classifier.invoke(_request())
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
assert exc_info.value.request_id == REQUEST_ID
|
||||
@@ -476,12 +526,11 @@ def test_connection_error_is_translated() -> None:
|
||||
client = httpx2.Client(transport=httpx2.MockTransport(handler))
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
client=client,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeSafeAPIConnectionError, match="Unable to connect"):
|
||||
classifier.invoke("hello")
|
||||
classifier.invoke(_request())
|
||||
|
||||
client.close()
|
||||
|
||||
@@ -499,12 +548,11 @@ def test_timeout_error_is_translated() -> None:
|
||||
)
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
client=client,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeSafeAPITimeoutError) as exc_info:
|
||||
classifier.invoke("hello")
|
||||
classifier.invoke(_request())
|
||||
|
||||
assert exc_info.value.timeout == client.timeout
|
||||
client.close()
|
||||
@@ -519,12 +567,11 @@ def test_invalid_response_is_translated() -> None:
|
||||
client = httpx2.Client(transport=httpx2.MockTransport(handler))
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
client=client,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeSafeAPIResponseValidationError, match="Invalid response"):
|
||||
classifier.invoke("hello")
|
||||
classifier.invoke(_request())
|
||||
|
||||
client.close()
|
||||
|
||||
@@ -549,11 +596,13 @@ def test_callbacks_receive_classifier_run() -> None:
|
||||
callback = RecordingHandler()
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions=_questions(),
|
||||
client=client,
|
||||
)
|
||||
|
||||
classifier.invoke("hello", config={"callbacks": [callback]})
|
||||
classifier.invoke(
|
||||
_request(),
|
||||
config={"callbacks": [callback]},
|
||||
)
|
||||
|
||||
assert callback.starts == 1
|
||||
assert callback.ends == 1
|
||||
@@ -578,14 +627,18 @@ def test_usage_is_recorded_on_the_active_run(
|
||||
recorder = _RunRecorder()
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions=_questions(),
|
||||
client=client,
|
||||
)
|
||||
|
||||
classifier.invoke("hello", config={"callbacks": [recorder]})
|
||||
classifier.invoke(
|
||||
_request(),
|
||||
config={"callbacks": [recorder]},
|
||||
)
|
||||
client.close()
|
||||
|
||||
assert recorder.run_type == "llm"
|
||||
assert recorder.input["state"] == "hello"
|
||||
assert set(recorder.input["questions"]) == set(_questions())
|
||||
assert stub.extra["metadata"]["usage_metadata"] == {
|
||||
"input_tokens": 42,
|
||||
"output_tokens": 12,
|
||||
@@ -606,11 +659,10 @@ async def test_async_usage_is_recorded_on_the_active_run(
|
||||
client = httpx2.AsyncClient(transport=httpx2.MockTransport(handler))
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions=_questions(),
|
||||
async_client=client,
|
||||
)
|
||||
|
||||
await classifier.ainvoke("hello")
|
||||
await classifier.ainvoke(_request())
|
||||
await client.aclose()
|
||||
|
||||
assert stub.extra["metadata"]["usage_metadata"]["total_tokens"] == 54
|
||||
@@ -626,12 +678,11 @@ def test_run_carries_model_identity_without_losing_caller_metadata() -> None:
|
||||
recorder = _RunRecorder()
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions=_questions(),
|
||||
client=client,
|
||||
)
|
||||
|
||||
classifier.invoke(
|
||||
"hello",
|
||||
_request(),
|
||||
config={"callbacks": [recorder], "metadata": {"tenant": "acme"}},
|
||||
)
|
||||
client.close()
|
||||
@@ -651,11 +702,10 @@ def test_untraced_invocation_is_unaffected() -> None:
|
||||
client = httpx2.Client(transport=httpx2.MockTransport(handler))
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions=_questions(),
|
||||
client=client,
|
||||
)
|
||||
|
||||
result = classifier.invoke("hello")
|
||||
result = classifier.invoke(_request())
|
||||
client.close()
|
||||
|
||||
assert result.usage.input_tokens == 42
|
||||
@@ -6,6 +6,8 @@ EXPECTED_ALL = [
|
||||
"Answer",
|
||||
"Choice",
|
||||
"ChoiceAnswer",
|
||||
"ClassifierRequest",
|
||||
"ClassifierResponse",
|
||||
"Noul",
|
||||
"NoulAnswer",
|
||||
"NoulCriteria",
|
||||
|
||||
Reference in new issue
Block a user