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:
Hunter Lovell authored and GitHub committed 2026-09-20 13:46:49 -05:00
1 parent f72f934cef
commit 115dbbd158
12 files changed
+385 -251

No files matched your search

+49 -24
View File
@@ -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,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",