diff --git a/libs/partners/typesafe/README.md b/libs/partners/typesafe/README.md index aa155d188a..794f4b86b1 100644 --- a/libs/partners/typesafe/README.md +++ b/libs/partners/typesafe/README.md @@ -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): diff --git a/libs/partners/typesafe/langchain_typesafe/__init__.py b/libs/partners/typesafe/langchain_typesafe/__init__.py index 4b9012e7a6..1d727853c4 100644 --- a/libs/partners/typesafe/langchain_typesafe/__init__.py +++ b/libs/partners/typesafe/langchain_typesafe/__init__.py @@ -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", diff --git a/libs/partners/typesafe/langchain_typesafe/classifier.py b/libs/partners/typesafe/langchain_typesafe/classifier.py index 6621afbfd0..0af17843ae 100644 --- a/libs/partners/typesafe/langchain_typesafe/classifier.py +++ b/libs/partners/typesafe/langchain_typesafe/classifier.py @@ -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() }, } diff --git a/libs/partners/typesafe/langchain_typesafe/client.py b/libs/partners/typesafe/langchain_typesafe/client.py index 15b3f0a4ca..5049a19ff1 100644 --- a/libs/partners/typesafe/langchain_typesafe/client.py +++ b/libs/partners/typesafe/langchain_typesafe/client.py @@ -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" diff --git a/libs/partners/typesafe/langchain_typesafe/experimental/middleware/auto_mode.py b/libs/partners/typesafe/langchain_typesafe/experimental/middleware/auto_mode.py index b53d1e7042..de1a8fce1a 100644 --- a/libs/partners/typesafe/langchain_typesafe/experimental/middleware/auto_mode.py +++ b/libs/partners/typesafe/langchain_typesafe/experimental/middleware/auto_mode.py @@ -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) diff --git a/libs/partners/typesafe/langchain_typesafe/experimental/middleware/model_router.py b/libs/partners/typesafe/langchain_typesafe/experimental/middleware/model_router.py index 2f431bed50..57e9ddbf17 100644 --- a/libs/partners/typesafe/langchain_typesafe/experimental/middleware/model_router.py +++ b/libs/partners/typesafe/langchain_typesafe/experimental/middleware/model_router.py @@ -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 diff --git a/libs/partners/typesafe/langchain_typesafe/types.py b/libs/partners/typesafe/langchain_typesafe/types.py index c05ca8c08f..3ba3f4b824 100644 --- a/libs/partners/typesafe/langchain_typesafe/types.py +++ b/libs/partners/typesafe/langchain_typesafe/types.py @@ -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", diff --git a/libs/partners/typesafe/tests/integration_tests/test_classifier.py b/libs/partners/typesafe/tests/integration_tests/test_classifier.py index 8ce1ec9658..0764d04074 100644 --- a/libs/partners/typesafe/tests/integration_tests/test_classifier.py +++ b/libs/partners/typesafe/tests/integration_tests/test_classifier.py @@ -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) diff --git a/libs/partners/typesafe/tests/unit_tests/experimental/middleware/test_auto_mode.py b/libs/partners/typesafe/tests/unit_tests/experimental/middleware/test_auto_mode.py index 9e2b02ef81..a57a90ac1d 100644 --- a/libs/partners/typesafe/tests/unit_tests/experimental/middleware/test_auto_mode.py +++ b/libs/partners/typesafe/tests/unit_tests/experimental/middleware/test_auto_mode.py @@ -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: diff --git a/libs/partners/typesafe/tests/unit_tests/experimental/middleware/test_model_router.py b/libs/partners/typesafe/tests/unit_tests/experimental/middleware/test_model_router.py index ef3aeda774..e822d41b5a 100644 --- a/libs/partners/typesafe/tests/unit_tests/experimental/middleware/test_model_router.py +++ b/libs/partners/typesafe/tests/unit_tests/experimental/middleware/test_model_router.py @@ -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"] diff --git a/libs/partners/typesafe/tests/unit_tests/test_classifier.py b/libs/partners/typesafe/tests/unit_tests/test_classifier.py index f63355a7ab..277853b821 100644 --- a/libs/partners/typesafe/tests/unit_tests/test_classifier.py +++ b/libs/partners/typesafe/tests/unit_tests/test_classifier.py @@ -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 diff --git a/libs/partners/typesafe/tests/unit_tests/test_imports.py b/libs/partners/typesafe/tests/unit_tests/test_imports.py index 7573ab7cc6..f3bb783482 100644 --- a/libs/partners/typesafe/tests/unit_tests/test_imports.py +++ b/libs/partners/typesafe/tests/unit_tests/test_imports.py @@ -6,6 +6,8 @@ EXPECTED_ALL = [ "Answer", "Choice", "ChoiceAnswer", + "ClassifierRequest", + "ClassifierResponse", "Noul", "NoulAnswer", "NoulCriteria",