mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 09:25:14 +03:00
feat(typesafe): TypeSafeClassifier (#40542)
Adds a first-party `langchain-typesafe` partner package for making
structured decisions using the new and shiny `jev` model. It introduces
a new `TypeSafeClassifier` runnable while using the typesafe wire
protocol, message normalization, response parsing, and provider failures
behind a standard LangChain interface.
- Supports TypeSafe's `Choice`, `Noul`, and `Score` primitives with
typed answers, probabilities, confidence, usage metadata, and request
IDs.
- Resolves credentials and API configuration from explicit arguments or
`TYPESAFE_API_KEY` and `TYPESAFE_BASE_URL`.
- Supports injected `httpx2.Client` and `httpx2.AsyncClient` instances
for custom transports, proxies, test fixtures, and shared connection
pools (matching patterns of other provider packages in this repo)
---
There were a couple of rapid fire design decisions that I'll document
here to aid in review:
### Extending base `Runnable`
This is largely the impetus for this PR; we can use typesafe's native
python client, but we lose tracing if it isn't captured through the
runnable interface. I imagine because this is a different model we want
to attribute its costs and things in tracing (at a later point in time)
### Direct `httpx2` integration
This package implements the System One `POST /v1/systemone` contract
directly rather than wrapping `typesafe-sdk` (the status quo of other
integrations in langchain). This keeps the integration's public types
and behavior native, avoids adding a transitive dependency, and gives us
lower level controls which is harder to do when indirecting through
someone elses client.
Aspirationally this is a direction we've wanted to work towards for
integrations generally:
* The benefit of wrapping an external sdk is that we can stay tied to
the supply-chain of a package thats maintained first party (we don't
need to maintain types, the transitive dependency keeps capabilities up
to date, etc.)
* This was in a time when maintainer attention was limited, and the
focus wasn't on nitty gritty implementation details with a provider
* ^Agents are probably good enough to start shouldering most of the
burden we've encountered in the past
* There's an indirection tax we have to pay when we shove the
responsibility onto a separate package- it makes us harder to design
single abstractions around
Because this is a 0-1 integration, my intention is to trial typesafes'
implementation under a "langchain maintained client" paradigm to see if
this is something that could work more generally
### State typing and message normalization
I had to bifurcate the input state types into `State` and `_StateValue`
representations mostly to respect TypeSafe API invariants: input must be
a string, object, or array. Its private recursive `_StateValue`
representation permits ordinary scalars by itself and through
dicts/sequences.
I'm also adding langchain `BaseMessage` objects and `BaseMessage`
sequences at any nesting level. This is because messages are the most
common unit of context within langchain, and translating between agent
context and classifier context is a very important part of us plugging
jev into the ecosystem. I'm imagining something like this:
```python
class ClassifierMiddleware(AgentMiddleware[AgentState, ContextT, ResponseT]):
def after_model(self, state):
result = TypeSafeClassifier(...).invoke(state.messages)
# do something with results
```
This uses the `convert_to_openai_messages` utility exposed in core since
that is the most context friendly representation we have to pass
messages through to typesafe (which doesn't have any concept of an LLM
message)
### Typed request and response models
Question and answer models use Pydantic because they sit directly on
serialization and validation layers, and I was hoping to meet some of
the invariant criteria that the typesafe API has. Things like:
* request models requiring instructions for every question
* `Score` requiring at least two ordered levels
* rejecting empty model identifiers
Criteria remain JSON-capable rather than being narrowed to strings
because TypeSafe's advanced documentation supports structured
instructions and criteria, even though the compact HTTP reference
presents narrower examples.
(lmk if pydantic is defunct and we should pivot)
### Error handling infrastructure
I wanted to meet the same level of specification that the TypeSafe
Python SDK has with regards to how it represents errors. This package
exports a standard error reference (in `_client_utils.py`) that does
this while also inheriting from alngchains standard `Model*Error`
classes which we recently added. Callers can handle either
TypeSafe-specific metadata or provider-independent model failures.
The package mirrors the TypeSafe Python SDK's status, connection,
timeout, rate-limit, and response-validation exception names.
Status-specific provider errors also inherit from LangChain's standard
`Model*Error` classes, following the OpenAI integration pattern, so
callers can handle either TypeSafe-specific metadata or
provider-independent model failures.
API errors retain structured status, body, headers, request ID,
sanitized endpoint, response-validation field path, and rate-limit delay
metadata where applicable. TypeSafe's `529 Overloaded` response maps to
a retryable `ModelAPIError`, while `429` maps to `ModelRateLimitError`
and preserves `retry_after_ms`.
Automatic retry policy is intentionally outside the initial package
scope. The classified retryable errors and retry metadata provide the
foundation for a focused follow-up without coupling this integration to
a retry policy in its first release.
## Package infrastructure
- Introduces `langchain-typesafe` at version `0.0.1` with typed-package
markers, documentation, unit and integration-compilation coverage, and a
reproducible `uv` lockfile.
- Adds bounded `httpx2>=2.0.0,<3.0.0` and `langchain-core>=1.6.2,<2.0.0`
dependencies. The core minimum provides the standard `Model*Error`
hierarchy used by the provider exceptions.
- Registers the package with repository CI, release automation,
dependency updates, issue routing, and package labeling.
## Release note
Added the new `langchain-typesafe` integration package, including a
composable `TypeSafeClassifier` for typed classification, scoring, and
confidence-aware decisions with TypeSafe's Jev model.
---
This was developed with heavy assistance from Sol, earnestly reviewed by
me, and validated through similar CI patterns established in other
packages
This commit is contained in:
1 parent
b3ec70925f
commit
d1fee60ef2
36 files changed
+5103
-1
No files matched your search
@@ -70,6 +70,7 @@ body:
|
||||
- label: langchain-openrouter
|
||||
- label: langchain-perplexity
|
||||
- label: langchain-qdrant
|
||||
- label: langchain-typesafe
|
||||
- label: langchain-xai
|
||||
- label: Other / not sure / general
|
||||
- type: textarea
|
||||
|
||||
@@ -69,6 +69,7 @@ body:
|
||||
- label: langchain-openrouter
|
||||
- label: langchain-perplexity
|
||||
- label: langchain-qdrant
|
||||
- label: langchain-typesafe
|
||||
- label: langchain-xai
|
||||
- label: Other / not sure / general
|
||||
- type: textarea
|
||||
|
||||
@@ -45,5 +45,6 @@ body:
|
||||
- label: langchain-openrouter
|
||||
- label: langchain-perplexity
|
||||
- label: langchain-qdrant
|
||||
- label: langchain-typesafe
|
||||
- label: langchain-xai
|
||||
- label: Other / not sure / general
|
||||
@@ -116,5 +116,6 @@ body:
|
||||
- label: langchain-openrouter
|
||||
- label: langchain-perplexity
|
||||
- label: langchain-qdrant
|
||||
- label: langchain-typesafe
|
||||
- label: langchain-xai
|
||||
- label: Other / not sure / general
|
||||
@@ -66,6 +66,7 @@ updates:
|
||||
- "/libs/partners/openrouter/"
|
||||
- "/libs/partners/perplexity/"
|
||||
- "/libs/partners/qdrant/"
|
||||
- "/libs/partners/typesafe/"
|
||||
- "/libs/partners/xai/"
|
||||
schedule:
|
||||
interval: "monthly"
|
||||
|
||||
@@ -47,6 +47,7 @@
|
||||
"openrouter": "openrouter",
|
||||
"perplexity": "perplexity",
|
||||
"qdrant": "qdrant",
|
||||
"typesafe": "typesafe",
|
||||
"xai": "xai",
|
||||
"deps": "dependencies",
|
||||
"docs": "documentation",
|
||||
@@ -74,6 +75,7 @@
|
||||
{ "label": "openrouter", "prefix": "libs/partners/openrouter/", "skipExcludedFiles": true },
|
||||
{ "label": "perplexity", "prefix": "libs/partners/perplexity/", "skipExcludedFiles": true },
|
||||
{ "label": "qdrant", "prefix": "libs/partners/qdrant/", "skipExcludedFiles": true },
|
||||
{ "label": "typesafe", "prefix": "libs/partners/typesafe/", "skipExcludedFiles": true },
|
||||
{ "label": "xai", "prefix": "libs/partners/xai/", "skipExcludedFiles": true },
|
||||
{ "label": "github_actions", "prefix": ".github/workflows/" },
|
||||
{ "label": "github_actions", "prefix": ".github/actions/" },
|
||||
|
||||
@@ -82,6 +82,7 @@ on:
|
||||
- openrouter
|
||||
- perplexity
|
||||
- qdrant
|
||||
- typesafe
|
||||
- xai
|
||||
working-directory-override:
|
||||
required: false
|
||||
|
||||
@@ -55,6 +55,7 @@ jobs:
|
||||
"langchain-openrouter": "openrouter",
|
||||
"langchain-perplexity": "perplexity",
|
||||
"langchain-qdrant": "qdrant",
|
||||
"langchain-typesafe": "typesafe",
|
||||
"langchain-xai": "xai",
|
||||
};
|
||||
|
||||
|
||||
@@ -43,6 +43,7 @@ on:
|
||||
- "openrouter"
|
||||
- "perplexity"
|
||||
- "qdrant"
|
||||
- "typesafe"
|
||||
- "xai"
|
||||
working-directory-override:
|
||||
type: string
|
||||
@@ -324,6 +325,7 @@ jobs:
|
||||
OPENROUTER_API_KEY: ${{ secrets.OPENROUTER_API_KEY }}
|
||||
PPLX_API_KEY: ${{ secrets.PPLX_API_KEY }}
|
||||
TOGETHER_API_KEY: ${{ secrets.TOGETHER_API_KEY }}
|
||||
TYPESAFE_API_KEY: ${{ secrets.TYPESAFE_API_KEY }}
|
||||
UPSTAGE_API_KEY: ${{ secrets.UPSTAGE_API_KEY }}
|
||||
WATSONX_APIKEY: ${{ secrets.WATSONX_APIKEY }}
|
||||
WATSONX_PROJECT_ID: ${{ secrets.WATSONX_PROJECT_ID }}
|
||||
|
||||
@@ -116,6 +116,7 @@ jobs:
|
||||
openrouter
|
||||
perplexity
|
||||
qdrant
|
||||
typesafe
|
||||
xai
|
||||
infra
|
||||
deps
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# Makefile for libs/partners/ directory
|
||||
# Contains targets that operate across all partner packages
|
||||
|
||||
PARTNER_DIRS = anthropic chroma deepseek exa fireworks groq huggingface mistralai nomic ollama openai openrouter perplexity qdrant xai
|
||||
PARTNER_DIRS = anthropic chroma deepseek exa fireworks groq huggingface mistralai nomic ollama openai openrouter perplexity qdrant typesafe xai
|
||||
|
||||
.PHONY: lock check-lock
|
||||
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
__pycache__
|
||||
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2026 LangChain, Inc.
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
@@ -0,0 +1,54 @@
|
||||
.PHONY: all format lint type test tests integration_tests help
|
||||
|
||||
all: help
|
||||
|
||||
.EXPORT_ALL_VARIABLES:
|
||||
UV_FROZEN = true
|
||||
|
||||
TEST_FILE ?= tests/unit_tests/
|
||||
PYTEST_EXTRA ?=
|
||||
integration_test integration_tests: TEST_FILE = tests/integration_tests/
|
||||
|
||||
test tests:
|
||||
env -u LANGCHAIN_TRACING_V2 -u LANGCHAIN_API_KEY -u LANGSMITH_API_KEY -u LANGSMITH_TRACING -u LANGCHAIN_PROJECT uv run --group test pytest $(PYTEST_EXTRA) --disable-socket --allow-unix-socket $(TEST_FILE)
|
||||
|
||||
integration_test integration_tests:
|
||||
uv run --group test --group test_integration pytest -v --tb=short -n auto $(PYTEST_EXTRA) $(TEST_FILE)
|
||||
|
||||
PYTHON_FILES=.
|
||||
MYPY_CACHE=.mypy_cache
|
||||
lint format: PYTHON_FILES=.
|
||||
lint_diff format_diff: PYTHON_FILES=$(shell git diff --relative=libs/partners/typesafe --name-only --diff-filter=d master | grep -E '\.py$$|\.ipynb$$')
|
||||
lint_package: PYTHON_FILES=langchain_typesafe
|
||||
lint_tests: PYTHON_FILES=tests
|
||||
lint_tests: MYPY_CACHE=.mypy_cache_test
|
||||
UV_RUN_LINT = uv run --all-groups
|
||||
UV_RUN_TYPE = uv run --all-groups
|
||||
lint_package lint_tests: UV_RUN_LINT = uv run --group lint
|
||||
|
||||
lint lint_diff lint_package lint_tests:
|
||||
./scripts/lint_imports.sh
|
||||
[ "$(PYTHON_FILES)" = "" ] || $(UV_RUN_LINT) ruff check $(PYTHON_FILES)
|
||||
[ "$(PYTHON_FILES)" = "" ] || $(UV_RUN_LINT) ruff format $(PYTHON_FILES) --diff
|
||||
[ "$(PYTHON_FILES)" = "" ] || mkdir -p $(MYPY_CACHE) && $(UV_RUN_TYPE) mypy $(PYTHON_FILES) --cache-dir $(MYPY_CACHE)
|
||||
|
||||
type:
|
||||
mkdir -p $(MYPY_CACHE) && $(UV_RUN_TYPE) mypy $(PYTHON_FILES) --cache-dir $(MYPY_CACHE)
|
||||
|
||||
format format_diff:
|
||||
[ "$(PYTHON_FILES)" = "" ] || $(UV_RUN_LINT) ruff format $(PYTHON_FILES)
|
||||
[ "$(PYTHON_FILES)" = "" ] || $(UV_RUN_LINT) ruff check --fix $(PYTHON_FILES)
|
||||
|
||||
check_imports: $(shell find langchain_typesafe -name '*.py')
|
||||
$(UV_RUN_LINT) python ./scripts/check_imports.py $^
|
||||
|
||||
check_version:
|
||||
uv run python ./scripts/check_version.py
|
||||
|
||||
help:
|
||||
@echo '----'
|
||||
@echo 'check_imports - check imports'
|
||||
@echo 'check_version - validate version consistency'
|
||||
@echo 'format - run code formatters'
|
||||
@echo 'lint - run linters and type checking'
|
||||
@echo 'test - run unit tests'
|
||||
@@ -0,0 +1,104 @@
|
||||
# langchain-typesafe
|
||||
|
||||
[](https://pypi.org/project/langchain-typesafe/#history)
|
||||
[](https://opensource.org/licenses/MIT)
|
||||
|
||||
## Installation
|
||||
|
||||
```bash
|
||||
uv add langchain-typesafe
|
||||
```
|
||||
|
||||
Set the `TYPESAFE_API_KEY` environment variable before making requests.
|
||||
|
||||
## Usage
|
||||
|
||||
`TypeSafeClassifier` is a LangChain `Runnable` for probabilistic classification and scoring with TypeSafe.
|
||||
|
||||
```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"],
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
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.
|
||||
|
||||
### LangChain messages as state
|
||||
|
||||
`BaseMessage` objects and message sequences can appear at the root or anywhere inside JSON state. The integration recursively converts them to objects with `role` and `content` fields while preserving surrounding application data:
|
||||
|
||||
```python
|
||||
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",
|
||||
}
|
||||
)
|
||||
```
|
||||
|
||||
### Custom HTTP clients
|
||||
|
||||
The classifier creates sync and async `httpx2` clients when they are not supplied. Applications that need custom transports, proxies, or shared connection pools can inject either client independently:
|
||||
|
||||
```python
|
||||
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"),
|
||||
)
|
||||
```
|
||||
|
||||
Injected clients are used as-is, and the application retains responsibility for their lifecycle.
|
||||
|
||||
### Error handling
|
||||
|
||||
Provider errors also inherit from LangChain's standard model-error hierarchy. Applications can therefore catch a TypeSafe-specific error when provider metadata is needed, or a LangChain error when handling several model providers uniformly:
|
||||
|
||||
```python
|
||||
from langchain_core.exceptions import ModelAuthenticationError, ModelRateLimitError
|
||||
from langchain_typesafe import TypeSafeRateLimitError
|
||||
|
||||
try:
|
||||
response = classifier.invoke("Classify this message.")
|
||||
except TypeSafeRateLimitError as error:
|
||||
print(error.request_id, error.retry_after_ms)
|
||||
except (ModelAuthenticationError, ModelRateLimitError):
|
||||
handle_model_error()
|
||||
```
|
||||
|
||||
`TypeSafeAPIError` exposes the response status, parsed body, headers, sanitized endpoint, and request ID. Connection, timeout, response-validation, and status-specific subclasses follow the names used by the TypeSafe Python SDK.
|
||||
|
||||
## Documentation
|
||||
|
||||
See the [TypeSafe documentation](https://docs.typesafe.ai/) for model and question semantics. LangChain API reference documentation is available at [reference.langchain.com](https://reference.langchain.com/python/integrations/langchain_typesafe/).
|
||||
|
||||
## Contributing
|
||||
|
||||
For contribution instructions, see the [LangChain contributing guide](https://docs.langchain.com/oss/python/contributing/overview).
|
||||
@@ -0,0 +1,33 @@
|
||||
"""LangChain integration for TypeSafe classifiers."""
|
||||
|
||||
from langchain_typesafe._version import __version__
|
||||
from langchain_typesafe.classifier import TypeSafeClassifier
|
||||
from langchain_typesafe.types import (
|
||||
Answer,
|
||||
Choice,
|
||||
ChoiceAnswer,
|
||||
Noul,
|
||||
NoulAnswer,
|
||||
NoulCriteria,
|
||||
Question,
|
||||
Score,
|
||||
ScoreAnswer,
|
||||
State,
|
||||
Usage,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"Answer",
|
||||
"Choice",
|
||||
"ChoiceAnswer",
|
||||
"Noul",
|
||||
"NoulAnswer",
|
||||
"NoulCriteria",
|
||||
"Question",
|
||||
"Score",
|
||||
"ScoreAnswer",
|
||||
"State",
|
||||
"TypeSafeClassifier",
|
||||
"Usage",
|
||||
"__version__",
|
||||
]
|
||||
@@ -0,0 +1,55 @@
|
||||
"""Normalize TypeSafe state containing LangChain messages."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from langchain_core.messages import BaseMessage, convert_to_openai_messages
|
||||
from pydantic import JsonValue
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from langchain_typesafe.types import State
|
||||
|
||||
|
||||
def _serialize_state_value(value: object) -> JsonValue:
|
||||
if isinstance(value, BaseMessage):
|
||||
return _serialize_state_value(convert_to_openai_messages(value))
|
||||
if isinstance(value, dict):
|
||||
if not all(isinstance(key, str) for key in value):
|
||||
message = "TypeSafe state object keys must be strings."
|
||||
raise TypeError(message)
|
||||
return {key: _serialize_state_value(item) for key, item in value.items()}
|
||||
if isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray)):
|
||||
return [_serialize_state_value(item) for item in value]
|
||||
if value is None or isinstance(value, (str, int, float, bool)):
|
||||
return value
|
||||
message = f"Unsupported TypeSafe state value: {type(value).__name__}."
|
||||
raise TypeError(message)
|
||||
|
||||
|
||||
def serialize_state(state: State) -> JsonValue:
|
||||
"""Recursively convert LangChain messages inside TypeSafe state to JSON.
|
||||
|
||||
Args:
|
||||
state: Native TypeSafe state containing zero or more LangChain messages.
|
||||
|
||||
Returns:
|
||||
A string, object, or array suitable for the TypeSafe `state` field. Every
|
||||
message or message sequence is replaced with role/content JSON while its
|
||||
surrounding object and array structure is preserved.
|
||||
|
||||
Raises:
|
||||
TypeError: If the root is a JSON scalar other than a string, or if any nested
|
||||
value cannot be represented as JSON or LangChain messages.
|
||||
"""
|
||||
if state is None or isinstance(state, (int, float, bool)):
|
||||
message = (
|
||||
"TypeSafe state must be a string, object, array, BaseMessage, or sequence "
|
||||
"of BaseMessage objects."
|
||||
)
|
||||
raise TypeError(message)
|
||||
return _serialize_state_value(state)
|
||||
|
||||
|
||||
__all__ = ["serialize_state"]
|
||||
@@ -0,0 +1,3 @@
|
||||
"""Version information for `langchain-typesafe`."""
|
||||
|
||||
__version__ = "0.0.1a1"
|
||||
@@ -0,0 +1,426 @@
|
||||
"""LangChain runnable for TypeSafe classification."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import httpx2
|
||||
from langchain_core._api import beta
|
||||
from langchain_core.runnables import RunnableConfig, RunnableSerializable
|
||||
from langchain_core.utils import from_env, secret_from_env
|
||||
from pydantic import (
|
||||
ConfigDict,
|
||||
Field,
|
||||
JsonValue,
|
||||
SecretStr,
|
||||
field_validator,
|
||||
model_validator,
|
||||
)
|
||||
from typing_extensions import Self, override
|
||||
|
||||
from langchain_typesafe._state import serialize_state
|
||||
from langchain_typesafe._version import __version__
|
||||
from langchain_typesafe.client import (
|
||||
TypeSafeAPIConnectionError,
|
||||
TypeSafeAPITimeoutError,
|
||||
parse_response,
|
||||
)
|
||||
from langchain_typesafe.types import ClassificationResponse, Question, State
|
||||
|
||||
_DEFAULT_BASE_URL = "https://api.typesafe.ai"
|
||||
_DEFAULT_MODEL = "jev-latest"
|
||||
_DEFAULT_TIMEOUT = 30.0
|
||||
|
||||
|
||||
@beta()
|
||||
class TypeSafeClassifier(RunnableSerializable[State, ClassificationResponse]):
|
||||
"""Classify JSON-compatible state with TypeSafe.
|
||||
|
||||
`TypeSafeClassifier` is a LangChain `Runnable` for asking one or more typed
|
||||
questions about text or structured state. A single request can combine binary
|
||||
`Noul` judgments, categorical `Choice` classifications, and ordinal `Score`
|
||||
evaluations. The response preserves probabilities, confidence, and token usage so
|
||||
application code can decide whether to act, route, or request human review.
|
||||
|
||||
After configuration is validated, the classifier creates both synchronous and
|
||||
asynchronous `httpx2` clients. Supply `client` and/or `async_client` to reuse
|
||||
clients configured by your application, including custom transports for testing,
|
||||
or network policy enforcement. Injected clients are used as-is; `timeout` only
|
||||
configures clients created by this class.
|
||||
|
||||
The classifier does not add separate client lifecycle methods. Keep classifier
|
||||
instances long-lived to benefit from connection pooling. Applications that require
|
||||
deterministic cleanup can close `classifier.client` and `classifier.async_client`
|
||||
directly, following the corresponding `httpx2` sync and async client interfaces.
|
||||
|
||||
Native TypeSafe state may be a string, JSON object, or JSON array. LangChain
|
||||
`BaseMessage` objects and message sequences can appear at the root or anywhere
|
||||
inside JSON objects and arrays. They are converted to role/content JSON before the
|
||||
request is sent. Message IDs are omitted, while system, user, assistant, and tool
|
||||
roles are preserved.
|
||||
|
||||
The API key is read from `TYPESAFE_API_KEY` when `api_key` is omitted. Explicit
|
||||
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.
|
||||
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.
|
||||
client: Optional synchronous `httpx2.Client` used by `invoke`.
|
||||
async_client: Optional asynchronous `httpx2.AsyncClient` used by `ainvoke`.
|
||||
|
||||
Raises:
|
||||
ValueError: If questions are empty, credentials are unavailable, or the timeout
|
||||
is not positive.
|
||||
|
||||
??? example "Classify state on several dimensions"
|
||||
|
||||
Send questions that share the same state together. Each answer remains
|
||||
independently addressable through its question ID.
|
||||
|
||||
```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."],
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
response = classifier.invoke(
|
||||
"Stripe has failed to connect for three days. Please help immediately."
|
||||
)
|
||||
print(response.choices["department"].choice)
|
||||
print(response.nouls["urgent"].noul)
|
||||
print(response.scores["frustration"].score)
|
||||
```
|
||||
|
||||
??? example "Classify asynchronously"
|
||||
|
||||
`ainvoke` uses the classifier's asynchronous HTTP client and returns the same
|
||||
`ClassificationResponse` type as `invoke`.
|
||||
|
||||
```python
|
||||
from langchain_typesafe import Noul, TypeSafeClassifier
|
||||
|
||||
classifier = TypeSafeClassifier(
|
||||
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.
|
||||
|
||||
The default, `jev-latest`, follows TypeSafe's latest compatible Jev release. Use a
|
||||
concrete model identifier when an application requires reproducible behavior across
|
||||
model updates. Leading and trailing whitespace is removed, and empty model names are
|
||||
rejected during initialization.
|
||||
"""
|
||||
|
||||
api_key: SecretStr | str = Field(
|
||||
default_factory=secret_from_env("TYPESAFE_API_KEY", default=""),
|
||||
exclude=True,
|
||||
repr=False,
|
||||
)
|
||||
"""API key used to authenticate TypeSafe requests.
|
||||
|
||||
If omitted, the key is read from the `TYPESAFE_API_KEY` environment variable when
|
||||
the classifier is initialized. An explicit constructor value takes precedence. The
|
||||
value is stored as `SecretStr` and excluded from model representation and
|
||||
serialization.
|
||||
|
||||
??? example "Specify with an environment variable"
|
||||
|
||||
```bash
|
||||
export TYPESAFE_API_KEY=...
|
||||
```
|
||||
|
||||
```python
|
||||
from langchain_typesafe import Noul, TypeSafeClassifier
|
||||
|
||||
classifier = TypeSafeClassifier(
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")}
|
||||
)
|
||||
```
|
||||
|
||||
??? example "Specify directly"
|
||||
|
||||
```python
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key="...",
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
)
|
||||
```
|
||||
"""
|
||||
|
||||
base_url: str = Field(
|
||||
default_factory=from_env("TYPESAFE_BASE_URL", default=_DEFAULT_BASE_URL)
|
||||
)
|
||||
"""Root URL used for TypeSafe API requests.
|
||||
|
||||
Resolution order:
|
||||
|
||||
1. Explicit `base_url` supplied to `TypeSafeClassifier`.
|
||||
2. The `TYPESAFE_BASE_URL` environment variable.
|
||||
3. `https://api.typesafe.ai`.
|
||||
|
||||
Requests are sent to `/v1/systemone` beneath this URL. Override it for a compatible
|
||||
gateway, test server, or private deployment. URL validation is delegated to
|
||||
`httpx2` when a request is made.
|
||||
"""
|
||||
|
||||
timeout: float = Field(default=_DEFAULT_TIMEOUT, gt=0)
|
||||
"""Timeout in seconds for clients created by this classifier.
|
||||
|
||||
This setting is passed to both `httpx2.Client` and `httpx2.AsyncClient` when their
|
||||
respective fields are omitted. It does not modify an injected client's timeout;
|
||||
configure custom clients directly when different sync and async policies are needed.
|
||||
"""
|
||||
|
||||
client: httpx2.Client | None = Field(default=None, exclude=True, repr=False)
|
||||
"""Optional synchronous `httpx2.Client` used by `invoke` and `batch`.
|
||||
|
||||
If omitted, the classifier creates a client using `timeout`. Supply a client to
|
||||
reuse connection pools or configure a custom transport, proxy, TLS policy, or test
|
||||
fixture. The injected client is used as-is and is not closed by the classifier; the
|
||||
caller retains responsibility for its lifecycle.
|
||||
|
||||
This client is not used by `ainvoke` or `abatch`. Configure `async_client`
|
||||
separately when asynchronous calls also require custom HTTP behavior.
|
||||
"""
|
||||
|
||||
async_client: httpx2.AsyncClient | None = Field(
|
||||
default=None,
|
||||
exclude=True,
|
||||
repr=False,
|
||||
)
|
||||
"""Optional asynchronous `httpx2.AsyncClient` used by `ainvoke` and `abatch`.
|
||||
|
||||
If omitted, the classifier creates an asynchronous client using `timeout`. Supply a
|
||||
client to reuse connection pools or configure a custom transport, proxy, TLS policy,
|
||||
or test fixture. The injected client is used as-is and is not closed by the
|
||||
classifier; the caller retains responsibility for its lifecycle.
|
||||
|
||||
This client is not used by `invoke` or `batch`. Configure `client` separately when
|
||||
synchronous calls also require custom HTTP behavior.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(
|
||||
arbitrary_types_allowed=True,
|
||||
extra="forbid",
|
||||
validate_default=True,
|
||||
)
|
||||
|
||||
@field_validator("model")
|
||||
@classmethod
|
||||
def _validate_model(cls, model: str) -> str:
|
||||
model = model.strip()
|
||||
if not model:
|
||||
message = "TypeSafe model must not be empty."
|
||||
raise ValueError(message)
|
||||
return model
|
||||
|
||||
@field_validator("api_key")
|
||||
@classmethod
|
||||
def _validate_api_key(cls, api_key: SecretStr | str) -> SecretStr:
|
||||
secret = api_key if isinstance(api_key, SecretStr) else SecretStr(api_key)
|
||||
if not secret.get_secret_value().strip():
|
||||
message = (
|
||||
"TypeSafe API key is required. Pass `api_key` or set "
|
||||
"`TYPESAFE_API_KEY`."
|
||||
)
|
||||
raise ValueError(message)
|
||||
return secret
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _build_clients(self) -> Self:
|
||||
"""Create missing sync and async clients after configuration is validated."""
|
||||
if self.client is None:
|
||||
self.client = httpx2.Client(timeout=self.timeout)
|
||||
if self.async_client is None:
|
||||
self.async_client = httpx2.AsyncClient(timeout=self.timeout)
|
||||
return self
|
||||
|
||||
@classmethod
|
||||
@override
|
||||
def is_lc_serializable(cls) -> bool:
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
@override
|
||||
def get_lc_namespace(cls) -> list[str]:
|
||||
return ["langchain", "classifiers", "typesafe"]
|
||||
|
||||
@property
|
||||
def lc_secrets(self) -> dict[str, str]:
|
||||
"""Map the API-key field to its environment variable for serialization."""
|
||||
return {"api_key": "TYPESAFE_API_KEY"}
|
||||
|
||||
@override
|
||||
def invoke(
|
||||
self,
|
||||
input: State,
|
||||
config: RunnableConfig | None = None,
|
||||
**_: Any,
|
||||
) -> ClassificationResponse:
|
||||
"""Classify one JSON-compatible input synchronously.
|
||||
|
||||
Args:
|
||||
input: Text, object, array, `BaseMessage`, or message sequence to classify.
|
||||
config: Optional LangChain runnable configuration for callbacks, tags,
|
||||
metadata, and tracing.
|
||||
**_: Additional keyword arguments accepted for `Runnable` compatibility and
|
||||
otherwise ignored.
|
||||
|
||||
Returns:
|
||||
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.
|
||||
TypeSafeAPIResponseValidationError: If a successful response is malformed.
|
||||
"""
|
||||
return self._call_with_config(
|
||||
self._classify,
|
||||
input,
|
||||
config,
|
||||
run_type="chain",
|
||||
)
|
||||
|
||||
@override
|
||||
async def ainvoke(
|
||||
self,
|
||||
input: State,
|
||||
config: RunnableConfig | None = None,
|
||||
**_: Any,
|
||||
) -> ClassificationResponse:
|
||||
"""Classify one JSON-compatible input asynchronously.
|
||||
|
||||
Args:
|
||||
input: Text, object, array, `BaseMessage`, or message sequence to classify.
|
||||
config: Optional LangChain runnable configuration for callbacks, tags,
|
||||
metadata, and tracing.
|
||||
**_: Additional keyword arguments accepted for `Runnable` compatibility and
|
||||
otherwise ignored.
|
||||
|
||||
Returns:
|
||||
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.
|
||||
TypeSafeAPIResponseValidationError: If a successful response is malformed.
|
||||
"""
|
||||
return await self._acall_with_config(
|
||||
self._aclassify,
|
||||
input,
|
||||
config,
|
||||
run_type="chain",
|
||||
)
|
||||
|
||||
def _classify(self, state: State) -> ClassificationResponse:
|
||||
payload = self._payload(state)
|
||||
if self.client is None: # pragma: no cover - guaranteed by model validation
|
||||
message = "Synchronous TypeSafe client was not initialized."
|
||||
raise TypeSafeAPIConnectionError(message)
|
||||
try:
|
||||
response = self.client.post(
|
||||
self._endpoint,
|
||||
json=payload,
|
||||
headers=self._request_headers,
|
||||
)
|
||||
except httpx2.TimeoutException as error:
|
||||
raise TypeSafeAPITimeoutError(self.client.timeout) from error
|
||||
except httpx2.HTTPError as error:
|
||||
message = "Unable to connect to the TypeSafe API."
|
||||
raise TypeSafeAPIConnectionError(message) from error
|
||||
return parse_response(response)
|
||||
|
||||
async def _aclassify(self, state: State) -> ClassificationResponse:
|
||||
payload = self._payload(state)
|
||||
if self.async_client is None: # pragma: no cover - guaranteed by validation
|
||||
message = "Asynchronous TypeSafe client was not initialized."
|
||||
raise TypeSafeAPIConnectionError(message)
|
||||
try:
|
||||
response = await self.async_client.post(
|
||||
self._endpoint,
|
||||
json=payload,
|
||||
headers=self._request_headers,
|
||||
)
|
||||
except httpx2.TimeoutException as error:
|
||||
raise TypeSafeAPITimeoutError(self.async_client.timeout) from error
|
||||
except httpx2.HTTPError as error:
|
||||
message = "Unable to connect to the TypeSafe API."
|
||||
raise TypeSafeAPIConnectionError(message) from error
|
||||
return parse_response(response)
|
||||
|
||||
@property
|
||||
def _endpoint(self) -> str:
|
||||
return f"{self.base_url.rstrip('/')}/v1/systemone"
|
||||
|
||||
@property
|
||||
def _request_headers(self) -> dict[str, str]:
|
||||
api_key = (
|
||||
self.api_key.get_secret_value()
|
||||
if isinstance(self.api_key, SecretStr)
|
||||
else self.api_key
|
||||
)
|
||||
return {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
"User-Agent": f"langchain-typesafe/{__version__}",
|
||||
}
|
||||
|
||||
def _payload(self, state: State) -> dict[str, JsonValue]:
|
||||
return {
|
||||
"state": serialize_state(state),
|
||||
"model": self.model,
|
||||
"questions": {
|
||||
name: question.model_dump(mode="json", exclude_none=True)
|
||||
for name, question in self.questions.items()
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
__all__ = ["TypeSafeClassifier"]
|
||||
@@ -0,0 +1,364 @@
|
||||
"""HTTP response handling and errors for the TypeSafe integration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import time
|
||||
from email.utils import parsedate_to_datetime
|
||||
from http import HTTPStatus
|
||||
from typing import Any
|
||||
from urllib.parse import urlsplit, urlunsplit
|
||||
|
||||
import httpx2
|
||||
from langchain_core.exceptions import (
|
||||
ModelAPIError,
|
||||
ModelAuthenticationError,
|
||||
ModelConnectionError,
|
||||
ModelInvalidRequestError,
|
||||
ModelNotFoundError,
|
||||
ModelPermissionDeniedError,
|
||||
ModelRateLimitError,
|
||||
ModelTimeoutError,
|
||||
)
|
||||
from pydantic import ValidationError
|
||||
from typing_extensions import override
|
||||
|
||||
from langchain_typesafe.types import ClassificationResponse
|
||||
|
||||
_REQUEST_ID_HEADER = "x-typesafe-request-id"
|
||||
_RETRY_AFTER_HEADER = "retry-after"
|
||||
_RETRY_AFTER_MS_HEADER = "retry-after-ms"
|
||||
_SAFE_STATUS_MESSAGES = {529: "Overloaded"}
|
||||
|
||||
|
||||
class TypeSafeError(Exception):
|
||||
"""Base exception for errors raised by the TypeSafe integration."""
|
||||
|
||||
|
||||
class TypeSafeAPIError(TypeSafeError):
|
||||
"""An unsuccessful HTTP response with its body and request metadata.
|
||||
|
||||
Attributes:
|
||||
status: HTTP response status code.
|
||||
body: Parsed JSON error body, plain response text, or `None`.
|
||||
headers: HTTP response headers.
|
||||
endpoint: Request method and URL without credentials, query, or fragment.
|
||||
request_id: Value of the `x-typesafe-request-id` response header.
|
||||
|
||||
The body and headers are available for programmatic error handling but deliberately
|
||||
excluded from `str` and `repr` to avoid exposing classified state or credentials in
|
||||
logs and tracebacks.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
status: int,
|
||||
body: Any,
|
||||
headers: httpx2.Headers,
|
||||
message: str | None = None,
|
||||
endpoint: str | None = None,
|
||||
) -> None:
|
||||
"""Create an API error from a TypeSafe HTTP response.
|
||||
|
||||
Args:
|
||||
status: HTTP response status code.
|
||||
body: Parsed JSON body, plain response text, or `None`.
|
||||
headers: HTTP response headers.
|
||||
message: Optional safe message override that does not contain response data.
|
||||
endpoint: Sanitized request method and URL, when available.
|
||||
"""
|
||||
super().__init__(status, body, headers, message, endpoint)
|
||||
self.status = status
|
||||
self.body = body
|
||||
self.headers = headers
|
||||
self.endpoint = endpoint
|
||||
self._message = message
|
||||
|
||||
@property
|
||||
def status_code(self) -> int:
|
||||
"""Alias for `status`, matching common HTTP exception interfaces."""
|
||||
return self.status
|
||||
|
||||
@property
|
||||
def request_id(self) -> str | None:
|
||||
"""Return the TypeSafe request ID from the response headers, when present."""
|
||||
return self.headers.get(_REQUEST_ID_HEADER)
|
||||
|
||||
@override
|
||||
def __str__(self) -> str:
|
||||
"""Describe the failure without including its response body or headers."""
|
||||
try:
|
||||
reason = HTTPStatus(self.status).phrase
|
||||
except ValueError:
|
||||
reason = _SAFE_STATUS_MESSAGES.get(self.status, "API request failed")
|
||||
detail = self._message or reason
|
||||
message = f"{self.status} {detail}"
|
||||
if self.endpoint is not None:
|
||||
message = f"{self.endpoint}: {message}"
|
||||
if self.request_id is not None:
|
||||
message = f"{message} (request_id={self.request_id})"
|
||||
return message
|
||||
|
||||
@override
|
||||
def __repr__(self) -> str:
|
||||
"""Represent the error without including its response body or headers."""
|
||||
return f"{type(self).__name__}({str(self)!r})"
|
||||
|
||||
|
||||
class TypeSafeBadRequestError(TypeSafeAPIError, ModelInvalidRequestError):
|
||||
"""The TypeSafe request was invalid (HTTP 400)."""
|
||||
|
||||
|
||||
class TypeSafeAuthenticationError(TypeSafeAPIError, ModelAuthenticationError):
|
||||
"""Authentication with TypeSafe failed (HTTP 401)."""
|
||||
|
||||
|
||||
class TypeSafePermissionDeniedError(TypeSafeAPIError, ModelPermissionDeniedError):
|
||||
"""The TypeSafe credential cannot perform the request (HTTP 403)."""
|
||||
|
||||
|
||||
class TypeSafeNotFoundError(TypeSafeAPIError, ModelNotFoundError):
|
||||
"""The requested TypeSafe resource or model was not found (HTTP 404)."""
|
||||
|
||||
|
||||
class TypeSafeUnprocessableEntityError(TypeSafeAPIError, ModelInvalidRequestError):
|
||||
"""TypeSafe rejected the request body during validation (HTTP 422)."""
|
||||
|
||||
|
||||
class TypeSafeRateLimitError(TypeSafeAPIError, ModelRateLimitError):
|
||||
"""The TypeSafe rate limit was exceeded (HTTP 429).
|
||||
|
||||
Attributes:
|
||||
retry_after_ms: Server-requested delay in milliseconds, or `None` when the
|
||||
response does not contain a valid retry header.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
status: int,
|
||||
body: Any,
|
||||
headers: httpx2.Headers,
|
||||
message: str | None = None,
|
||||
endpoint: str | None = None,
|
||||
) -> None:
|
||||
"""Create a rate-limit error and parse its retry delay.
|
||||
|
||||
Args:
|
||||
status: HTTP response status code.
|
||||
body: Parsed JSON body, plain response text, or `None`.
|
||||
headers: HTTP response headers.
|
||||
message: Optional safe message override.
|
||||
endpoint: Sanitized request method and URL, when available.
|
||||
"""
|
||||
super().__init__(status, body, headers, message, endpoint)
|
||||
self.retry_after_ms = _parse_retry_after(headers)
|
||||
|
||||
|
||||
class TypeSafeInternalServerError(TypeSafeAPIError, ModelAPIError):
|
||||
"""TypeSafe failed to process the request (HTTP 5xx)."""
|
||||
|
||||
|
||||
class TypeSafeAPIConnectionError(TypeSafeError, ModelConnectionError, ConnectionError):
|
||||
"""A TypeSafe request failed without receiving an HTTP response."""
|
||||
|
||||
|
||||
class TypeSafeAPITimeoutError(
|
||||
TypeSafeAPIConnectionError,
|
||||
ModelTimeoutError,
|
||||
TimeoutError,
|
||||
):
|
||||
"""A TypeSafe request exceeded its configured timeout.
|
||||
|
||||
Attributes:
|
||||
timeout: Timeout setting used by the sync or async HTTP client.
|
||||
"""
|
||||
|
||||
def __init__(self, timeout: float | httpx2.Timeout) -> None:
|
||||
"""Create a timeout error.
|
||||
|
||||
Args:
|
||||
timeout: Timeout setting used for the failed request.
|
||||
"""
|
||||
super().__init__(timeout)
|
||||
self.timeout = timeout
|
||||
|
||||
@override
|
||||
def __str__(self) -> str:
|
||||
"""Return the configured timeout without request or credential data."""
|
||||
return f"Request timed out (timeout={self.timeout})."
|
||||
|
||||
@override
|
||||
def __repr__(self) -> str:
|
||||
"""Represent the timeout using its safe formatted message."""
|
||||
return f"{type(self).__name__}({str(self)!r})"
|
||||
|
||||
|
||||
class TypeSafeAPIResponseValidationError(TypeSafeAPIError):
|
||||
"""A successful response was missing or contained invalid required data.
|
||||
|
||||
Attributes:
|
||||
field_path: Dotted path to the first field that failed validation.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
status: int,
|
||||
body: Any,
|
||||
headers: httpx2.Headers,
|
||||
field_path: str,
|
||||
endpoint: str | None = None,
|
||||
) -> None:
|
||||
"""Create a response-validation error.
|
||||
|
||||
Args:
|
||||
status: Successful HTTP response status code.
|
||||
body: Parsed JSON body or plain response text.
|
||||
headers: HTTP response headers.
|
||||
field_path: Dotted path to the first invalid field.
|
||||
endpoint: Sanitized request method and URL, when available.
|
||||
"""
|
||||
self.field_path = field_path
|
||||
super().__init__(
|
||||
status,
|
||||
body,
|
||||
headers,
|
||||
f"Invalid response data at {field_path!r}.",
|
||||
endpoint,
|
||||
)
|
||||
self.args = (status, body, headers, field_path, endpoint)
|
||||
|
||||
|
||||
def _parse_retry_after(headers: httpx2.Headers) -> float | None:
|
||||
for name, multiplier in (
|
||||
(_RETRY_AFTER_MS_HEADER, 1.0),
|
||||
(_RETRY_AFTER_HEADER, 1000.0),
|
||||
):
|
||||
raw = headers.get(name)
|
||||
if raw is None:
|
||||
continue
|
||||
try:
|
||||
value = float(raw.strip() or "0")
|
||||
except ValueError:
|
||||
if name == _RETRY_AFTER_HEADER:
|
||||
try:
|
||||
delay = (
|
||||
parsedate_to_datetime(raw).timestamp() - time.time()
|
||||
) * 1000
|
||||
except (OverflowError, TypeError, ValueError):
|
||||
continue
|
||||
return max(0.0, delay)
|
||||
continue
|
||||
if math.isfinite(value) and value >= 0:
|
||||
delay = value * multiplier
|
||||
if math.isfinite(delay):
|
||||
return delay
|
||||
if name == _RETRY_AFTER_HEADER:
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def _response_body(response: httpx2.Response) -> Any:
|
||||
if not response.content:
|
||||
return None
|
||||
try:
|
||||
return response.json()
|
||||
except ValueError:
|
||||
return response.text
|
||||
|
||||
|
||||
def _response_endpoint(response: httpx2.Response) -> str | None:
|
||||
try:
|
||||
request = response.request
|
||||
except RuntimeError:
|
||||
return None
|
||||
parts = urlsplit(str(request.url))
|
||||
hostname = parts.hostname
|
||||
if hostname is None:
|
||||
return None
|
||||
host = f"[{hostname}]" if ":" in hostname else hostname
|
||||
try:
|
||||
port = parts.port
|
||||
except ValueError:
|
||||
port = None
|
||||
netloc = f"{host}:{port}" if port is not None else host
|
||||
url = urlunsplit((parts.scheme, netloc, parts.path, "", ""))
|
||||
return f"{request.method} {url}"
|
||||
|
||||
|
||||
_STATUS_ERROR_TYPES: dict[int, type[TypeSafeAPIError]] = {
|
||||
400: TypeSafeBadRequestError,
|
||||
401: TypeSafeAuthenticationError,
|
||||
403: TypeSafePermissionDeniedError,
|
||||
404: TypeSafeNotFoundError,
|
||||
422: TypeSafeUnprocessableEntityError,
|
||||
429: TypeSafeRateLimitError,
|
||||
}
|
||||
|
||||
|
||||
def _api_error(response: httpx2.Response) -> TypeSafeAPIError:
|
||||
error_type = _STATUS_ERROR_TYPES.get(
|
||||
response.status_code,
|
||||
TypeSafeInternalServerError
|
||||
if response.status_code >= HTTPStatus.INTERNAL_SERVER_ERROR
|
||||
else TypeSafeAPIError,
|
||||
)
|
||||
return error_type(
|
||||
response.status_code,
|
||||
_response_body(response),
|
||||
response.headers,
|
||||
endpoint=_response_endpoint(response),
|
||||
)
|
||||
|
||||
|
||||
def parse_response(response: httpx2.Response) -> ClassificationResponse:
|
||||
"""Validate an HTTP response and convert it to a classification response.
|
||||
|
||||
Args:
|
||||
response: Raw HTTP response returned by the TypeSafe API.
|
||||
|
||||
Returns:
|
||||
Validated classification answers and metadata. The TypeSafe request ID is
|
||||
copied from the response headers when present.
|
||||
|
||||
Raises:
|
||||
TypeSafeAPIError: If TypeSafe returns an unsuccessful status code. Specific
|
||||
statuses use subclasses that also inherit from LangChain model errors.
|
||||
TypeSafeAPIResponseValidationError: If a successful response is not valid JSON
|
||||
or does not match the expected response schema.
|
||||
"""
|
||||
if not response.is_success:
|
||||
raise _api_error(response)
|
||||
endpoint = _response_endpoint(response)
|
||||
body = _response_body(response)
|
||||
try:
|
||||
parsed = ClassificationResponse.model_validate(body)
|
||||
except ValidationError as error:
|
||||
location = error.errors()[0].get("loc", ())
|
||||
field_path = ".".join(str(item) for item in location) or "response"
|
||||
raise TypeSafeAPIResponseValidationError(
|
||||
response.status_code,
|
||||
body,
|
||||
response.headers,
|
||||
field_path,
|
||||
endpoint,
|
||||
) from error
|
||||
return parsed.model_copy(
|
||||
update={"request_id": response.headers.get(_REQUEST_ID_HEADER)}
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"TypeSafeAPIConnectionError",
|
||||
"TypeSafeAPIError",
|
||||
"TypeSafeAPIResponseValidationError",
|
||||
"TypeSafeAPITimeoutError",
|
||||
"TypeSafeAuthenticationError",
|
||||
"TypeSafeBadRequestError",
|
||||
"TypeSafeError",
|
||||
"TypeSafeInternalServerError",
|
||||
"TypeSafeNotFoundError",
|
||||
"TypeSafePermissionDeniedError",
|
||||
"TypeSafeRateLimitError",
|
||||
"TypeSafeUnprocessableEntityError",
|
||||
"parse_response",
|
||||
]
|
||||
Whitespace-only changes.
@@ -0,0 +1,324 @@
|
||||
"""Question and response types for the TypeSafe integration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from typing import Annotated, Literal, TypeAlias
|
||||
|
||||
from langchain_core.messages import BaseMessage
|
||||
from pydantic import BaseModel, ConfigDict, Field, JsonValue
|
||||
|
||||
_QuestionContent: TypeAlias = str | dict[str, JsonValue] | list[JsonValue]
|
||||
_StateValue: TypeAlias = (
|
||||
str
|
||||
| int
|
||||
| float
|
||||
| bool
|
||||
| BaseMessage
|
||||
| Sequence["_StateValue"]
|
||||
| dict[str, "_StateValue"]
|
||||
| None
|
||||
)
|
||||
|
||||
State: TypeAlias = str | BaseMessage | Sequence[_StateValue] | dict[str, _StateValue]
|
||||
"""Root state accepted by `TypeSafeClassifier`.
|
||||
|
||||
TypeSafe natively accepts a string, JSON object, or JSON array. LangChain
|
||||
`BaseMessage` objects and message sequences can appear at the root or at any depth
|
||||
inside objects and arrays. The integration serializes messages as role/content JSON
|
||||
while preserving surrounding JSON structure.
|
||||
"""
|
||||
|
||||
|
||||
class NoulCriteria(BaseModel):
|
||||
"""Optional descriptions for the two outcomes of a `Noul` question.
|
||||
|
||||
Criteria clarify what should count as yes and no when the instruction alone leaves
|
||||
room for interpretation. Both values accept any JSON-compatible content, so callers
|
||||
can provide a short description or structured examples.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
|
||||
true: JsonValue = None
|
||||
"""Description of the yes outcome, or `None` when no clarification is needed."""
|
||||
|
||||
false: JsonValue = None
|
||||
"""Description of the no outcome, or `None` when no clarification is needed."""
|
||||
|
||||
|
||||
class Noul(BaseModel):
|
||||
"""Ask a binary question and receive the probability that its answer is yes.
|
||||
|
||||
Use `Noul` when the probability itself is useful to application code, such as
|
||||
deciding whether a message reports a bug or requests a refund. A value near `1`
|
||||
indicates strong support for yes, a value near `0` indicates strong support for no,
|
||||
and a value near `0.5` indicates uncertainty. Noul answers do not include a separate
|
||||
confidence value.
|
||||
|
||||
??? example "Detect an urgent support request"
|
||||
|
||||
```python
|
||||
from langchain_typesafe import Noul, TypeSafeClassifier
|
||||
|
||||
classifier = TypeSafeClassifier(
|
||||
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:
|
||||
page_on_call_engineer()
|
||||
```
|
||||
"""
|
||||
|
||||
type: Literal["noul"] = "noul"
|
||||
"""Wire discriminator for a binary TypeSafe question."""
|
||||
|
||||
instructions: _QuestionContent
|
||||
"""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.
|
||||
"""
|
||||
|
||||
criteria: NoulCriteria | None = None
|
||||
"""Optional descriptions that define what the yes and no outcomes mean."""
|
||||
|
||||
|
||||
class Choice(BaseModel):
|
||||
"""Select one label from a fixed set of alternatives.
|
||||
|
||||
Use `Choice` for categorical decisions with no inherent ordering, such as routing a
|
||||
support request, detecting a document type, or selecting an intent. The answer
|
||||
includes the selected label, a probability for every supplied label, and confidence
|
||||
derived from the shape of that probability distribution.
|
||||
|
||||
??? example "Route a support request"
|
||||
|
||||
```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.",
|
||||
},
|
||||
)
|
||||
}
|
||||
)
|
||||
response = classifier.invoke("Stripe fails whenever I connect my account.")
|
||||
department = response.choices["department"]
|
||||
|
||||
if department.confidence >= 0.7:
|
||||
route_to(department.choice)
|
||||
else:
|
||||
route_to_human_triage()
|
||||
```
|
||||
"""
|
||||
|
||||
type: Literal["choice"] = "choice"
|
||||
"""Wire discriminator for a categorical TypeSafe question."""
|
||||
|
||||
criteria: dict[str, JsonValue] = Field(min_length=1)
|
||||
"""Candidate labels mapped to their descriptions.
|
||||
|
||||
Descriptions may be text, structured JSON, or `None`. Include an `other` or
|
||||
`none_of_the_above` label when the supplied alternatives may not cover every input.
|
||||
"""
|
||||
|
||||
instructions: _QuestionContent
|
||||
"""Complete categorical judgment to make about the input state."""
|
||||
|
||||
|
||||
class Score(BaseModel):
|
||||
"""Evaluate state against an ordered rubric.
|
||||
|
||||
Use `Score` when the answer lies on a spectrum whose levels can be described, such
|
||||
as severity, urgency, or customer frustration. Criteria are numbered from zero in
|
||||
their supplied order. The returned score is an expected value and may fall between
|
||||
integer levels; the full probability distribution remains available for custom
|
||||
decision logic.
|
||||
|
||||
??? example "Score customer frustration"
|
||||
|
||||
```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.",
|
||||
],
|
||||
)
|
||||
}
|
||||
)
|
||||
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.
|
||||
print(frustration.legend) # The original zero-based rubric.
|
||||
print(frustration.probabilities) # Probability for each rubric level.
|
||||
```
|
||||
"""
|
||||
|
||||
type: Literal["score"] = "score"
|
||||
"""Wire discriminator for an ordinal TypeSafe question."""
|
||||
|
||||
criteria: list[JsonValue] = Field(min_length=2)
|
||||
"""Two or more ordered descriptions for score levels starting at zero."""
|
||||
|
||||
instructions: _QuestionContent
|
||||
"""Complete ordinal judgment to make about the input state."""
|
||||
|
||||
|
||||
Question = Annotated[Noul | Choice | Score, Field(discriminator="type")]
|
||||
"""A discriminated union of question types accepted by `TypeSafeClassifier`."""
|
||||
|
||||
|
||||
class NoulAnswer(BaseModel):
|
||||
"""Probability that a `Noul` question's answer is yes."""
|
||||
|
||||
type: Literal["noul"]
|
||||
"""Wire discriminator identifying a binary answer."""
|
||||
|
||||
noul: float = Field(ge=0.0, le=1.0)
|
||||
"""Probability of yes in the inclusive range from `0` to `1`."""
|
||||
|
||||
|
||||
class ChoiceAnswer(BaseModel):
|
||||
"""Selected `Choice` label with its probability distribution and confidence."""
|
||||
|
||||
type: Literal["choice"]
|
||||
"""Wire discriminator identifying a categorical answer."""
|
||||
|
||||
choice: str
|
||||
"""Label selected from the options supplied in `Choice.criteria`."""
|
||||
|
||||
probabilities: dict[str, float]
|
||||
"""Probability assigned to each candidate label, keyed by label name.
|
||||
|
||||
The complete distribution is retained so applications can use a confidence measure
|
||||
or risk policy different from TypeSafe's default confidence calculation.
|
||||
"""
|
||||
|
||||
confidence: float = Field(ge=0.0, le=1.0)
|
||||
"""Scalar certainty from `0` to `1`, derived from the probability distribution.
|
||||
|
||||
Confidence describes how concentrated the distribution is; it is not the selected
|
||||
label's probability. Thresholds should be chosen according to the consequences of
|
||||
an incorrect automated decision.
|
||||
"""
|
||||
|
||||
|
||||
class ScoreAnswer(BaseModel):
|
||||
"""Expected `Score` value with its rubric, distribution, and confidence."""
|
||||
|
||||
type: Literal["score"]
|
||||
"""Wire discriminator identifying an ordinal answer."""
|
||||
|
||||
score: float
|
||||
"""Expected position on the ordered rubric.
|
||||
|
||||
This value may be fractional because it summarizes the probability distribution
|
||||
over integer rubric levels rather than selecting exactly one level.
|
||||
"""
|
||||
|
||||
legend: dict[int, JsonValue]
|
||||
"""Original rubric descriptions keyed by their zero-based integer levels."""
|
||||
|
||||
probabilities: dict[int, float]
|
||||
"""Probability distribution over the zero-based integer rubric levels."""
|
||||
|
||||
confidence: float = Field(ge=0.0, le=1.0)
|
||||
"""Scalar certainty from `0` to `1`, derived from the level distribution."""
|
||||
|
||||
|
||||
Answer = Annotated[NoulAnswer | ChoiceAnswer | ScoreAnswer, Field(discriminator="type")]
|
||||
"""A discriminated union of answers returned by `TypeSafeClassifier`."""
|
||||
|
||||
|
||||
class Usage(BaseModel):
|
||||
"""Token usage reported for a TypeSafe classification request."""
|
||||
|
||||
input_tokens: int | None = None
|
||||
"""Number of input tokens processed, or `None` when not reported."""
|
||||
|
||||
output_tokens: int | None = None
|
||||
"""Number of output tokens produced, or `None` when not reported."""
|
||||
|
||||
|
||||
class ClassificationResponse(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`.
|
||||
"""
|
||||
|
||||
model: str
|
||||
"""TypeSafe model that answered the request."""
|
||||
|
||||
answers: dict[str, Answer]
|
||||
"""All recognized answers keyed by their original question IDs."""
|
||||
|
||||
usage: Usage = Field(default_factory=Usage)
|
||||
"""Input and output token counts reported for the request."""
|
||||
|
||||
request_id: str | None = None
|
||||
"""TypeSafe request ID, useful when diagnosing a request with provider support."""
|
||||
|
||||
@property
|
||||
def nouls(self) -> dict[str, NoulAnswer]:
|
||||
"""Return binary answers keyed by their original question IDs."""
|
||||
return {
|
||||
name: answer
|
||||
for name, answer in self.answers.items()
|
||||
if isinstance(answer, NoulAnswer)
|
||||
}
|
||||
|
||||
@property
|
||||
def choices(self) -> dict[str, ChoiceAnswer]:
|
||||
"""Return categorical answers keyed by their original question IDs."""
|
||||
return {
|
||||
name: answer
|
||||
for name, answer in self.answers.items()
|
||||
if isinstance(answer, ChoiceAnswer)
|
||||
}
|
||||
|
||||
@property
|
||||
def scores(self) -> dict[str, ScoreAnswer]:
|
||||
"""Return ordinal answers keyed by their original question IDs."""
|
||||
return {
|
||||
name: answer
|
||||
for name, answer in self.answers.items()
|
||||
if isinstance(answer, ScoreAnswer)
|
||||
}
|
||||
|
||||
|
||||
__all__ = [
|
||||
"Answer",
|
||||
"Choice",
|
||||
"ChoiceAnswer",
|
||||
"ClassificationResponse",
|
||||
"Noul",
|
||||
"NoulAnswer",
|
||||
"NoulCriteria",
|
||||
"Question",
|
||||
"Score",
|
||||
"ScoreAnswer",
|
||||
"State",
|
||||
"Usage",
|
||||
]
|
||||
@@ -0,0 +1,102 @@
|
||||
[build-system]
|
||||
requires = ["hatchling"]
|
||||
build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "langchain-typesafe"
|
||||
version = "0.0.1a1"
|
||||
description = "A LangChain integration for TypeSafe classifiers"
|
||||
license = { text = "MIT" }
|
||||
readme = "README.md"
|
||||
classifiers = [
|
||||
"Development Status :: 4 - Beta",
|
||||
"Intended Audience :: Developers",
|
||||
"License :: OSI Approved :: MIT License",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3.10",
|
||||
"Programming Language :: Python :: 3.11",
|
||||
"Programming Language :: Python :: 3.12",
|
||||
"Programming Language :: Python :: 3.13",
|
||||
"Programming Language :: Python :: 3.14",
|
||||
"Topic :: Scientific/Engineering :: Artificial Intelligence",
|
||||
]
|
||||
requires-python = ">=3.10.0,<4.0.0"
|
||||
dependencies = [
|
||||
"httpx2>=2.0.0,<3.0.0",
|
||||
"langchain-core>=1.6.2,<2.0.0",
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
Homepage = "https://docs.langchain.com/oss/python/integrations/providers/typesafe"
|
||||
Documentation = "https://reference.langchain.com/python/integrations/langchain_typesafe/"
|
||||
Repository = "https://github.com/langchain-ai/langchain"
|
||||
Issues = "https://github.com/langchain-ai/langchain/issues"
|
||||
Changelog = "https://github.com/langchain-ai/langchain/releases?q=%22langchain-typesafe%22"
|
||||
Twitter = "https://x.com/langchain_oss"
|
||||
Slack = "https://www.langchain.com/join-community"
|
||||
Reddit = "https://www.reddit.com/r/LangChain/"
|
||||
|
||||
[dependency-groups]
|
||||
test = [
|
||||
"pytest>=9.0.3,<10.0.0",
|
||||
"pytest-asyncio>=1.3.0,<2.0.0",
|
||||
"pytest-socket>=0.7.0,<1.0.0",
|
||||
"pytest-watcher>=0.6.3,<1.0.0",
|
||||
"pytest-xdist>=3.6.1,<4.0.0",
|
||||
"langchain-tests>=1.1.9,<2.0.0",
|
||||
]
|
||||
test_integration = []
|
||||
lint = ["ruff>=0.15.0,<0.16.0"]
|
||||
dev = []
|
||||
typing = ["mypy>=2.1.0,<2.2.0"]
|
||||
|
||||
[tool.uv]
|
||||
constraint-dependencies = ["pygments>=2.20.0"] # CVE-2026-4539
|
||||
|
||||
[tool.uv.sources]
|
||||
langchain-core = { path = "../../core", editable = true }
|
||||
langchain-tests = { path = "../../standard-tests", editable = true }
|
||||
|
||||
[tool.mypy]
|
||||
disallow_untyped_defs = true
|
||||
|
||||
[tool.ruff.format]
|
||||
docstring-code-format = true
|
||||
|
||||
[tool.ruff.lint]
|
||||
select = ["ALL"]
|
||||
ignore = [
|
||||
"COM812", # Conflicts with formatter
|
||||
"PLR0913", # Too many arguments
|
||||
"ANN401", # Any is required at LangChain callback seams
|
||||
"TC002", # Runtime imports are useful for Pydantic
|
||||
"TC003", # Runtime imports are useful for Pydantic
|
||||
]
|
||||
unfixable = ["B028"] # People should intentionally tune the stacklevel
|
||||
|
||||
[tool.ruff.lint.pydocstyle]
|
||||
convention = "google"
|
||||
ignore-var-parameters = true
|
||||
|
||||
[tool.ruff.lint.flake8-tidy-imports]
|
||||
ban-relative-imports = "all"
|
||||
|
||||
[tool.coverage.run]
|
||||
omit = ["tests/*"]
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
addopts = "--strict-markers --strict-config --durations=5"
|
||||
markers = [
|
||||
"compile: mark placeholder test used to compile integration tests without running them",
|
||||
]
|
||||
asyncio_mode = "auto"
|
||||
|
||||
[tool.ruff.lint.extend-per-file-ignores]
|
||||
"tests/**/*.py" = [
|
||||
"S101", # Tests need assertions
|
||||
"SLF001", # Private member access
|
||||
"PLR2004", # Magic values are fine in tests
|
||||
]
|
||||
"scripts/*.py" = [
|
||||
"INP001", # Not a package
|
||||
]
|
||||
@@ -0,0 +1,19 @@
|
||||
"""Script to check imports of given Python files."""
|
||||
|
||||
import sys
|
||||
import traceback
|
||||
from importlib.machinery import SourceFileLoader
|
||||
|
||||
if __name__ == "__main__":
|
||||
files = sys.argv[1:]
|
||||
has_failure = False
|
||||
for file in files:
|
||||
try:
|
||||
SourceFileLoader("x", file).load_module()
|
||||
except Exception: # noqa: PERF203, BLE001
|
||||
has_failure = True
|
||||
print(file) # noqa: T201
|
||||
traceback.print_exc()
|
||||
print() # noqa: T201
|
||||
|
||||
sys.exit(1 if has_failure else 0)
|
||||
@@ -0,0 +1,33 @@
|
||||
"""Check `langchain-typesafe` version consistency."""
|
||||
|
||||
import re
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _read_version(path: Path, pattern: str) -> str | None:
|
||||
content = path.read_text(encoding="utf-8")
|
||||
match = re.search(pattern, content, re.MULTILINE)
|
||||
return match.group(1) if match else None
|
||||
|
||||
|
||||
def main() -> int:
|
||||
"""Return a nonzero status when package versions differ."""
|
||||
package_dir = Path(__file__).parent.parent
|
||||
pyproject_version = _read_version(
|
||||
package_dir / "pyproject.toml",
|
||||
r'^version\s*=\s*"([^"]+)"',
|
||||
)
|
||||
module_version = _read_version(
|
||||
package_dir / "langchain_typesafe" / "_version.py",
|
||||
r'^__version__\s*=\s*"([^"]+)"',
|
||||
)
|
||||
if pyproject_version != module_version or pyproject_version is None:
|
||||
print("Error: package versions do not match.") # noqa: T201
|
||||
return 1
|
||||
print(f"Version check passed: {pyproject_version}") # noqa: T201
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
+12
@@ -0,0 +1,12 @@
|
||||
#!/bin/bash
|
||||
|
||||
set -eu
|
||||
|
||||
errors=0
|
||||
|
||||
git --no-pager grep "^from langchain\." . | grep -v ":from langchain\.agents" | grep -v ":from langchain\.tools" && errors=$((errors+1))
|
||||
git --no-pager grep "^from langchain_experimental\." . && errors=$((errors+1))
|
||||
|
||||
if [ "$errors" -gt 0 ]; then
|
||||
exit 1
|
||||
fi
|
||||
@@ -0,0 +1 @@
|
||||
"""Tests for `langchain-typesafe`."""
|
||||
@@ -0,0 +1 @@
|
||||
"""Integration tests for `langchain-typesafe`."""
|
||||
@@ -0,0 +1,126 @@
|
||||
"""Live integration tests for `TypeSafeClassifier`."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import HumanMessage, SystemMessage
|
||||
|
||||
from langchain_typesafe import (
|
||||
Choice,
|
||||
ChoiceAnswer,
|
||||
Noul,
|
||||
NoulAnswer,
|
||||
Score,
|
||||
ScoreAnswer,
|
||||
TypeSafeClassifier,
|
||||
)
|
||||
|
||||
|
||||
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.",
|
||||
],
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
try:
|
||||
response = classifier.invoke(
|
||||
{
|
||||
"message": (
|
||||
"Stripe has failed to connect for three days. "
|
||||
"Please help immediately."
|
||||
),
|
||||
"account_tier": "enterprise",
|
||||
}
|
||||
)
|
||||
|
||||
department = response.answers["department"]
|
||||
urgent = response.answers["urgent"]
|
||||
frustration = response.answers["frustration"]
|
||||
|
||||
assert isinstance(department, ChoiceAnswer)
|
||||
assert department.choice in labels
|
||||
assert set(department.probabilities) == labels
|
||||
assert sum(department.probabilities.values()) == pytest.approx(1.0)
|
||||
|
||||
assert isinstance(urgent, NoulAnswer)
|
||||
assert 0 <= urgent.noul <= 1
|
||||
|
||||
assert isinstance(frustration, ScoreAnswer)
|
||||
assert 0 <= frustration.score <= 2
|
||||
assert set(frustration.legend) == {0, 1, 2}
|
||||
assert set(frustration.probabilities) == {0, 1, 2}
|
||||
assert sum(frustration.probabilities.values()) == pytest.approx(1.0)
|
||||
|
||||
assert response.model.startswith("jev-")
|
||||
assert response.request_id
|
||||
assert isinstance(response.usage.input_tokens, int)
|
||||
assert isinstance(response.usage.output_tokens, int)
|
||||
finally:
|
||||
if classifier.client is not None:
|
||||
classifier.client.close()
|
||||
if classifier.async_client is not None:
|
||||
asyncio.run(classifier.async_client.aclose())
|
||||
|
||||
|
||||
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?"
|
||||
)
|
||||
}
|
||||
)
|
||||
|
||||
try:
|
||||
response = await classifier.ainvoke(
|
||||
{
|
||||
"conversation": [
|
||||
SystemMessage("You are reviewing a customer support conversation."),
|
||||
HumanMessage(
|
||||
"The integration crashes every time I connect Stripe. "
|
||||
"Can someone help?"
|
||||
),
|
||||
],
|
||||
"account": {
|
||||
"tier": "enterprise",
|
||||
"failed_attempts": 3,
|
||||
"trial": False,
|
||||
"notes": None,
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
needs_support = response.answers["needs_support"]
|
||||
assert isinstance(needs_support, NoulAnswer)
|
||||
assert 0 <= needs_support.noul <= 1
|
||||
assert response.model.startswith("jev-")
|
||||
assert response.request_id
|
||||
finally:
|
||||
if classifier.async_client is not None:
|
||||
await classifier.async_client.aclose()
|
||||
if classifier.client is not None:
|
||||
classifier.client.close()
|
||||
@@ -0,0 +1,8 @@
|
||||
"""Test compilation of integration tests."""
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.mark.compile
|
||||
def test_placeholder() -> None:
|
||||
"""Provide a target for the integration-test compilation job."""
|
||||
@@ -0,0 +1 @@
|
||||
"""Unit tests for `langchain-typesafe`."""
|
||||
@@ -0,0 +1,539 @@
|
||||
"""Unit tests for `TypeSafeClassifier`."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
import httpx2
|
||||
import pytest
|
||||
from langchain_core._api import LangChainBetaWarning
|
||||
from langchain_core.callbacks import BaseCallbackHandler
|
||||
from langchain_core.messages import AIMessage, HumanMessage, SystemMessage
|
||||
from pydantic import SecretStr, ValidationError
|
||||
|
||||
from langchain_typesafe import (
|
||||
Choice,
|
||||
ChoiceAnswer,
|
||||
Noul,
|
||||
NoulAnswer,
|
||||
Score,
|
||||
ScoreAnswer,
|
||||
TypeSafeClassifier,
|
||||
__version__,
|
||||
)
|
||||
from langchain_typesafe.client import (
|
||||
TypeSafeAPIConnectionError,
|
||||
TypeSafeAPIError,
|
||||
TypeSafeAPIResponseValidationError,
|
||||
TypeSafeAPITimeoutError,
|
||||
)
|
||||
|
||||
API_KEY = "test-api-key"
|
||||
REQUEST_ID = "req_test"
|
||||
|
||||
|
||||
def _response_payload() -> dict[str, Any]:
|
||||
return {
|
||||
"model": "jev-latest",
|
||||
"answers": {
|
||||
"department": {
|
||||
"type": "choice",
|
||||
"choice": "technical",
|
||||
"probabilities": {"billing": 0.1, "technical": 0.9},
|
||||
"confidence": 0.8,
|
||||
},
|
||||
"urgent": {"type": "noul", "noul": 0.95},
|
||||
"frustration": {
|
||||
"type": "score",
|
||||
"score": 1.25,
|
||||
"legend": {"0": "calm", "1": "frustrated", "2": "angry"},
|
||||
"probabilities": {"0": 0.1, "1": 0.55, "2": 0.35},
|
||||
"confidence": 0.7,
|
||||
},
|
||||
},
|
||||
"usage": {"input_tokens": 42, "output_tokens": 12},
|
||||
}
|
||||
|
||||
|
||||
def _questions() -> dict[str, Choice | Noul | Score]:
|
||||
return {
|
||||
"department": Choice(
|
||||
instructions="Which team should handle this?",
|
||||
criteria={"billing": "Payment issues", "technical": None},
|
||||
),
|
||||
"urgent": Noul(instructions="Is this urgent?"),
|
||||
"frustration": Score(
|
||||
instructions="How frustrated is the customer?",
|
||||
criteria=["calm", "frustrated", "angry"],
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def test_classifier_is_beta() -> None:
|
||||
"""Constructing the classifier warns that its API is in beta."""
|
||||
with pytest.warns(
|
||||
LangChainBetaWarning,
|
||||
match=r"The class `TypeSafeClassifier` is in beta\.",
|
||||
):
|
||||
TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
)
|
||||
|
||||
|
||||
def test_questions_require_instructions() -> None:
|
||||
"""Every TypeSafe question requires an explicit instruction."""
|
||||
with pytest.raises(ValidationError, match="instructions"):
|
||||
Noul.model_validate({})
|
||||
with pytest.raises(ValidationError, match="instructions"):
|
||||
Choice.model_validate({"criteria": {"billing": None}})
|
||||
with pytest.raises(ValidationError, match="instructions"):
|
||||
Score.model_validate({"criteria": ["low", "high"]})
|
||||
|
||||
|
||||
def test_score_requires_two_levels() -> None:
|
||||
"""A Score rubric must define at least two ordered levels."""
|
||||
with pytest.raises(ValidationError, match="at least 2"):
|
||||
Score(instructions="How urgent is this?", criteria=["low"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", ["", " "])
|
||||
def test_model_must_not_be_empty(model: str) -> None:
|
||||
"""The classifier rejects empty and whitespace-only model identifiers."""
|
||||
with pytest.raises(ValidationError):
|
||||
TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
model=model,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
)
|
||||
|
||||
|
||||
def test_invoke_sends_request_and_parses_response() -> None:
|
||||
"""The sync runnable sends the expected wire payload and parses each answer."""
|
||||
|
||||
def handler(request: httpx2.Request) -> httpx2.Response:
|
||||
assert request.url == "https://api.typesafe.ai/v1/systemone"
|
||||
assert request.headers["authorization"] == f"Bearer {API_KEY}"
|
||||
assert request.headers["user-agent"] == f"langchain-typesafe/{__version__}"
|
||||
payload = json.loads(request.content)
|
||||
assert payload == {
|
||||
"state": {"message": "Stripe fails to connect."},
|
||||
"model": "jev-latest",
|
||||
"questions": {
|
||||
"department": {
|
||||
"type": "choice",
|
||||
"criteria": {
|
||||
"billing": "Payment issues",
|
||||
"technical": None,
|
||||
},
|
||||
"instructions": "Which team should handle this?",
|
||||
},
|
||||
"urgent": {
|
||||
"type": "noul",
|
||||
"instructions": "Is this urgent?",
|
||||
},
|
||||
"frustration": {
|
||||
"type": "score",
|
||||
"criteria": ["calm", "frustrated", "angry"],
|
||||
"instructions": "How frustrated is the customer?",
|
||||
},
|
||||
},
|
||||
}
|
||||
return httpx2.Response(
|
||||
200,
|
||||
json=_response_payload(),
|
||||
headers={"x-typesafe-request-id": REQUEST_ID},
|
||||
)
|
||||
|
||||
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."})
|
||||
|
||||
assert result.request_id == REQUEST_ID
|
||||
assert result.usage.input_tokens == 42
|
||||
assert result.choices["department"] == ChoiceAnswer(
|
||||
type="choice",
|
||||
choice="technical",
|
||||
probabilities={"billing": 0.1, "technical": 0.9},
|
||||
confidence=0.8,
|
||||
)
|
||||
assert result.nouls["urgent"] == NoulAnswer(type="noul", noul=0.95)
|
||||
assert result.scores["frustration"] == ScoreAnswer(
|
||||
type="score",
|
||||
score=1.25,
|
||||
legend={0: "calm", 1: "frustrated", 2: "angry"},
|
||||
probabilities={0: 0.1, 1: 0.55, 2: 0.35},
|
||||
confidence=0.7,
|
||||
)
|
||||
client.close()
|
||||
|
||||
|
||||
def test_single_message_is_serialized_as_role_content_state() -> None:
|
||||
"""A LangChain message becomes a Jev-friendly role/content object."""
|
||||
observed_state: Any = None
|
||||
|
||||
def handler(request: httpx2.Request) -> httpx2.Response:
|
||||
nonlocal observed_state
|
||||
observed_state = json.loads(request.content)["state"]
|
||||
return httpx2.Response(200, json=_response_payload())
|
||||
|
||||
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."))
|
||||
|
||||
assert observed_state == {
|
||||
"role": "user",
|
||||
"content": "Please help immediately.",
|
||||
}
|
||||
client.close()
|
||||
|
||||
|
||||
def test_message_sequence_is_serialized_as_conversation_state() -> None:
|
||||
"""A message sequence preserves system, user, and assistant roles for Jev."""
|
||||
observed_state: Any = None
|
||||
|
||||
def handler(request: httpx2.Request) -> httpx2.Response:
|
||||
nonlocal observed_state
|
||||
observed_state = json.loads(request.content)["state"]
|
||||
return httpx2.Response(200, json=_response_payload())
|
||||
|
||||
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."),
|
||||
]
|
||||
)
|
||||
|
||||
assert observed_state == [
|
||||
{"role": "system", "content": "You are a support assistant."},
|
||||
{"role": "user", "content": "My integration is broken."},
|
||||
{"role": "assistant", "content": "I can help troubleshoot it."},
|
||||
]
|
||||
client.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ainvoke_uses_async_client() -> None:
|
||||
"""The async runnable sends requests through the injected async client."""
|
||||
|
||||
async def handler(request: httpx2.Request) -> httpx2.Response:
|
||||
assert request.headers["authorization"] == f"Bearer {API_KEY}"
|
||||
return httpx2.Response(200, json=_response_payload())
|
||||
|
||||
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.")
|
||||
|
||||
assert result.choices["department"].choice == "technical"
|
||||
await async_client.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_clients_are_created(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Initialization creates both clients with the configured timeout when omitted."""
|
||||
sync_client = httpx2.Client()
|
||||
async_client = httpx2.AsyncClient()
|
||||
observed_timeouts: list[float] = []
|
||||
|
||||
def sync_factory(*, timeout: float) -> httpx2.Client:
|
||||
observed_timeouts.append(timeout)
|
||||
return sync_client
|
||||
|
||||
def async_factory(*, timeout: float) -> httpx2.AsyncClient:
|
||||
observed_timeouts.append(timeout)
|
||||
return async_client
|
||||
|
||||
monkeypatch.setattr(httpx2, "Client", sync_factory)
|
||||
monkeypatch.setattr(httpx2, "AsyncClient", async_factory)
|
||||
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
timeout=12.5,
|
||||
)
|
||||
|
||||
assert classifier.client is sync_client
|
||||
assert classifier.async_client is async_client
|
||||
assert observed_timeouts == [12.5, 12.5]
|
||||
sync_client.close()
|
||||
await async_client.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_injected_clients_are_preserved() -> None:
|
||||
"""Initialization does not replace clients configured by the caller."""
|
||||
client = httpx2.Client()
|
||||
async_client = httpx2.AsyncClient()
|
||||
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
client=client,
|
||||
async_client=async_client,
|
||||
)
|
||||
|
||||
assert classifier.client is client
|
||||
assert classifier.async_client is async_client
|
||||
client.close()
|
||||
await async_client.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ainvoke_translates_api_error() -> None:
|
||||
"""The async runnable translates unsuccessful API responses."""
|
||||
|
||||
async def handler(_: httpx2.Request) -> httpx2.Response:
|
||||
return httpx2.Response(
|
||||
429,
|
||||
headers={"x-typesafe-request-id": REQUEST_ID},
|
||||
)
|
||||
|
||||
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")
|
||||
|
||||
assert exc_info.value.status_code == 429
|
||||
assert exc_info.value.request_id == REQUEST_ID
|
||||
await async_client.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ainvoke_translates_connection_error() -> None:
|
||||
"""The async runnable translates HTTP transport failures."""
|
||||
|
||||
async def handler(request: httpx2.Request) -> httpx2.Response:
|
||||
message = "sensitive transport detail"
|
||||
raise httpx2.ConnectError(message, request=request)
|
||||
|
||||
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 async_client.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ainvoke_translates_timeout_error() -> None:
|
||||
"""The async runnable classifies HTTP timeouts separately from connections."""
|
||||
|
||||
async def handler(request: httpx2.Request) -> httpx2.Response:
|
||||
message = "request timed out"
|
||||
raise httpx2.ReadTimeout(message, request=request)
|
||||
|
||||
async_client = httpx2.AsyncClient(
|
||||
timeout=6.0,
|
||||
transport=httpx2.MockTransport(handler),
|
||||
)
|
||||
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")
|
||||
|
||||
assert exc_info.value.timeout == async_client.timeout
|
||||
await async_client.aclose()
|
||||
|
||||
|
||||
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?")}
|
||||
)
|
||||
assert isinstance(classifier.api_key, SecretStr)
|
||||
assert classifier.api_key.get_secret_value() == API_KEY
|
||||
|
||||
|
||||
def test_base_url_from_environment(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""The classifier reads its API root from `TYPESAFE_BASE_URL`."""
|
||||
monkeypatch.setenv("TYPESAFE_BASE_URL", "https://gateway.typesafe.example")
|
||||
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions={"urgent": Noul(instructions="Is this urgent?")},
|
||||
)
|
||||
|
||||
assert classifier.base_url == "https://gateway.typesafe.example"
|
||||
|
||||
|
||||
def test_explicit_base_url_overrides_environment(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""An explicit API root takes precedence over environment configuration."""
|
||||
monkeypatch.setenv("TYPESAFE_BASE_URL", "https://environment.example")
|
||||
|
||||
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"
|
||||
|
||||
|
||||
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?")})
|
||||
|
||||
|
||||
def test_api_error_does_not_expose_response_body() -> None:
|
||||
"""HTTP errors expose status and request ID, but not server response bodies."""
|
||||
|
||||
def handler(_: httpx2.Request) -> httpx2.Response:
|
||||
return httpx2.Response(
|
||||
401,
|
||||
json={"error": "secret diagnostic"},
|
||||
headers={"x-typesafe-request-id": REQUEST_ID},
|
||||
)
|
||||
|
||||
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")
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
assert exc_info.value.request_id == REQUEST_ID
|
||||
assert "secret diagnostic" not in str(exc_info.value)
|
||||
client.close()
|
||||
|
||||
|
||||
def test_connection_error_is_translated() -> None:
|
||||
"""HTTP transport failures use the package exception hierarchy."""
|
||||
|
||||
def handler(request: httpx2.Request) -> httpx2.Response:
|
||||
message = "sensitive transport detail"
|
||||
raise httpx2.ConnectError(message, request=request)
|
||||
|
||||
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")
|
||||
|
||||
client.close()
|
||||
|
||||
|
||||
def test_timeout_error_is_translated() -> None:
|
||||
"""HTTP timeout failures use provider and LangChain timeout hierarchies."""
|
||||
|
||||
def handler(request: httpx2.Request) -> httpx2.Response:
|
||||
message = "request timed out"
|
||||
raise httpx2.ReadTimeout(message, request=request)
|
||||
|
||||
client = httpx2.Client(
|
||||
timeout=7.5,
|
||||
transport=httpx2.MockTransport(handler),
|
||||
)
|
||||
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")
|
||||
|
||||
assert exc_info.value.timeout == client.timeout
|
||||
client.close()
|
||||
|
||||
|
||||
def test_invalid_response_is_translated() -> None:
|
||||
"""Malformed successful responses raise a stable package exception."""
|
||||
|
||||
def handler(_: httpx2.Request) -> httpx2.Response:
|
||||
return httpx2.Response(200, json={"model": "jev-latest", "answers": []})
|
||||
|
||||
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")
|
||||
|
||||
client.close()
|
||||
|
||||
|
||||
def test_callbacks_receive_classifier_run() -> None:
|
||||
"""Invocation participates in the standard LangChain callback lifecycle."""
|
||||
|
||||
class RecordingHandler(BaseCallbackHandler):
|
||||
starts = 0
|
||||
ends = 0
|
||||
|
||||
def on_chain_start(self, *_: Any, **__: Any) -> None:
|
||||
self.starts += 1
|
||||
|
||||
def on_chain_end(self, *_: Any, **__: Any) -> None:
|
||||
self.ends += 1
|
||||
|
||||
def handler(_: httpx2.Request) -> httpx2.Response:
|
||||
return httpx2.Response(200, json=_response_payload())
|
||||
|
||||
client = httpx2.Client(transport=httpx2.MockTransport(handler))
|
||||
callback = RecordingHandler()
|
||||
classifier = TypeSafeClassifier(
|
||||
api_key=API_KEY,
|
||||
questions=_questions(),
|
||||
client=client,
|
||||
)
|
||||
|
||||
classifier.invoke("hello", config={"callbacks": [callback]})
|
||||
|
||||
assert callback.starts == 1
|
||||
assert callback.ends == 1
|
||||
client.close()
|
||||
@@ -0,0 +1,215 @@
|
||||
"""Tests for TypeSafe provider and LangChain error classification."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pickle
|
||||
from typing import Any
|
||||
|
||||
import httpx2
|
||||
import pytest
|
||||
from langchain_core.exceptions import (
|
||||
ModelAPIError,
|
||||
ModelAuthenticationError,
|
||||
ModelConnectionError,
|
||||
ModelError,
|
||||
ModelInvalidRequestError,
|
||||
ModelNotFoundError,
|
||||
ModelPermissionDeniedError,
|
||||
ModelRateLimitError,
|
||||
ModelTimeoutError,
|
||||
)
|
||||
|
||||
from langchain_typesafe.client import (
|
||||
TypeSafeAPIConnectionError,
|
||||
TypeSafeAPIError,
|
||||
TypeSafeAPIResponseValidationError,
|
||||
TypeSafeAPITimeoutError,
|
||||
TypeSafeAuthenticationError,
|
||||
TypeSafeBadRequestError,
|
||||
TypeSafeInternalServerError,
|
||||
TypeSafeNotFoundError,
|
||||
TypeSafePermissionDeniedError,
|
||||
TypeSafeRateLimitError,
|
||||
TypeSafeUnprocessableEntityError,
|
||||
parse_response,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("status", "provider_type", "langchain_type", "is_retryable"),
|
||||
[
|
||||
(400, TypeSafeBadRequestError, ModelInvalidRequestError, False),
|
||||
(401, TypeSafeAuthenticationError, ModelAuthenticationError, False),
|
||||
(403, TypeSafePermissionDeniedError, ModelPermissionDeniedError, False),
|
||||
(404, TypeSafeNotFoundError, ModelNotFoundError, False),
|
||||
(422, TypeSafeUnprocessableEntityError, ModelInvalidRequestError, False),
|
||||
(429, TypeSafeRateLimitError, ModelRateLimitError, True),
|
||||
(500, TypeSafeInternalServerError, ModelAPIError, True),
|
||||
(529, TypeSafeInternalServerError, ModelAPIError, True),
|
||||
],
|
||||
)
|
||||
def test_status_errors_use_provider_and_langchain_types(
|
||||
status: int,
|
||||
provider_type: type[TypeSafeAPIError],
|
||||
langchain_type: type[ModelError],
|
||||
*,
|
||||
is_retryable: bool,
|
||||
) -> None:
|
||||
"""Each known status is catchable through provider and LangChain hierarchies."""
|
||||
request = httpx2.Request("POST", "https://api.typesafe.ai/v1/systemone")
|
||||
response = httpx2.Response(status, json={"message": "failure"}, request=request)
|
||||
|
||||
with pytest.raises(provider_type) as exc_info:
|
||||
parse_response(response)
|
||||
|
||||
assert isinstance(exc_info.value, langchain_type)
|
||||
assert exc_info.value.is_retryable is is_retryable
|
||||
|
||||
|
||||
def test_overloaded_error_has_safe_provider_description() -> None:
|
||||
"""TypeSafe's nonstandard overloaded status remains useful without body text."""
|
||||
response = httpx2.Response(
|
||||
529,
|
||||
json={"message": "private overload detail"},
|
||||
request=httpx2.Request("POST", "https://api.typesafe.ai/v1/systemone"),
|
||||
)
|
||||
|
||||
with pytest.raises(TypeSafeInternalServerError) as exc_info:
|
||||
parse_response(response)
|
||||
|
||||
assert "529 Overloaded" in str(exc_info.value)
|
||||
assert "private overload detail" not in str(exc_info.value)
|
||||
|
||||
|
||||
def test_api_error_exposes_metadata_without_leaking_it_in_repr() -> None:
|
||||
"""API errors expose structured context while keeping string forms sanitized."""
|
||||
request = httpx2.Request(
|
||||
"POST",
|
||||
"https://user:password@example.test/v1/systemone?token=secret#fragment",
|
||||
)
|
||||
response = httpx2.Response(
|
||||
400,
|
||||
json={"message": "private response detail"},
|
||||
headers={"x-typesafe-request-id": "req_123"},
|
||||
request=request,
|
||||
)
|
||||
|
||||
with pytest.raises(TypeSafeBadRequestError) as exc_info:
|
||||
parse_response(response)
|
||||
|
||||
error = exc_info.value
|
||||
assert error.status == 400
|
||||
assert error.status_code == 400
|
||||
assert error.body == {"message": "private response detail"}
|
||||
assert error.headers["x-typesafe-request-id"] == "req_123"
|
||||
assert error.request_id == "req_123"
|
||||
assert error.endpoint == "POST https://example.test/v1/systemone"
|
||||
assert "private response detail" not in str(error)
|
||||
assert "password" not in repr(error)
|
||||
assert "token=secret" not in repr(error)
|
||||
|
||||
|
||||
def test_endpoint_sanitization_preserves_ipv6_and_port() -> None:
|
||||
"""Sanitization retains IPv6 addressing and explicit ports."""
|
||||
request = httpx2.Request(
|
||||
"POST",
|
||||
"https://[2001:db8::1]:8443/v1/systemone?token=secret",
|
||||
)
|
||||
response = httpx2.Response(400, request=request)
|
||||
|
||||
with pytest.raises(TypeSafeBadRequestError) as exc_info:
|
||||
parse_response(response)
|
||||
|
||||
assert exc_info.value.endpoint == "POST https://[2001:db8::1]:8443/v1/systemone"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("headers", "expected"),
|
||||
[
|
||||
({"retry-after-ms": "125"}, 125.0),
|
||||
({"retry-after": "2"}, 2000.0),
|
||||
({"retry-after-ms": "bad", "retry-after": "3"}, 3000.0),
|
||||
({"retry-after-ms": "bad", "retry-after": "bad"}, None),
|
||||
],
|
||||
)
|
||||
def test_rate_limit_error_parses_retry_delay(
|
||||
headers: dict[str, str], expected: float | None
|
||||
) -> None:
|
||||
"""Rate-limit responses expose the server-requested delay in milliseconds."""
|
||||
response = httpx2.Response(
|
||||
429,
|
||||
headers=headers,
|
||||
request=httpx2.Request("POST", "https://api.typesafe.ai/v1/systemone"),
|
||||
)
|
||||
|
||||
with pytest.raises(TypeSafeRateLimitError) as exc_info:
|
||||
parse_response(response)
|
||||
|
||||
assert exc_info.value.retry_after_ms == expected
|
||||
|
||||
|
||||
def test_response_validation_error_reports_field_path() -> None:
|
||||
"""Malformed successful responses identify the first invalid field."""
|
||||
body: dict[str, Any] = {"model": "jev-latest", "answers": []}
|
||||
response = httpx2.Response(
|
||||
200,
|
||||
json=body,
|
||||
request=httpx2.Request("POST", "https://api.typesafe.ai/v1/systemone"),
|
||||
)
|
||||
|
||||
with pytest.raises(TypeSafeAPIResponseValidationError) as exc_info:
|
||||
parse_response(response)
|
||||
|
||||
error = exc_info.value
|
||||
assert error.status == 200
|
||||
assert error.body == body
|
||||
assert error.field_path == "answers"
|
||||
|
||||
|
||||
def test_connection_error_uses_standard_hierarchies() -> None:
|
||||
"""Connection failures are catchable as provider, LangChain, and Python errors."""
|
||||
error = TypeSafeAPIConnectionError("Unable to connect")
|
||||
|
||||
assert isinstance(error, ModelConnectionError)
|
||||
assert isinstance(error, ConnectionError)
|
||||
assert error.is_retryable is True
|
||||
|
||||
|
||||
def test_timeout_error_uses_standard_hierarchies() -> None:
|
||||
"""Timeouts retain their setting and all provider and standard base types."""
|
||||
timeout = httpx2.Timeout(10.0)
|
||||
error = TypeSafeAPITimeoutError(timeout)
|
||||
|
||||
assert isinstance(error, TypeSafeAPIConnectionError)
|
||||
assert isinstance(error, ModelTimeoutError)
|
||||
assert isinstance(error, TimeoutError)
|
||||
assert error.is_retryable is True
|
||||
assert error.timeout is timeout
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"error",
|
||||
[
|
||||
TypeSafeAuthenticationError(401, {}, httpx2.Headers()),
|
||||
TypeSafeRateLimitError(
|
||||
429,
|
||||
{},
|
||||
httpx2.Headers({"retry-after-ms": "125"}),
|
||||
),
|
||||
TypeSafeAPITimeoutError(10.0),
|
||||
TypeSafeAPIResponseValidationError(
|
||||
200,
|
||||
{},
|
||||
httpx2.Headers(),
|
||||
"answers.urgent.noul",
|
||||
),
|
||||
],
|
||||
ids=lambda error: type(error).__name__,
|
||||
)
|
||||
def test_errors_round_trip_through_pickle(error: Exception) -> None:
|
||||
"""Structured errors retain their type and attributes across process boundaries."""
|
||||
restored = pickle.loads(pickle.dumps(error)) # noqa: S301
|
||||
|
||||
assert type(restored) is type(error)
|
||||
assert restored.args == error.args
|
||||
assert vars(restored) == vars(error)
|
||||
@@ -0,0 +1,24 @@
|
||||
"""Test the `langchain_typesafe` public interface."""
|
||||
|
||||
from langchain_typesafe import __all__
|
||||
|
||||
EXPECTED_ALL = [
|
||||
"Answer",
|
||||
"Choice",
|
||||
"ChoiceAnswer",
|
||||
"Noul",
|
||||
"NoulAnswer",
|
||||
"NoulCriteria",
|
||||
"Question",
|
||||
"Score",
|
||||
"ScoreAnswer",
|
||||
"State",
|
||||
"TypeSafeClassifier",
|
||||
"Usage",
|
||||
"__version__",
|
||||
]
|
||||
|
||||
|
||||
def test_all_imports() -> None:
|
||||
"""Verify that `__all__` contains the intended public interface."""
|
||||
assert sorted(EXPECTED_ALL) == sorted(__all__)
|
||||
@@ -0,0 +1,108 @@
|
||||
"""Tests for TypeSafe state and LangChain message normalization."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import AIMessage, HumanMessage, SystemMessage, ToolMessage
|
||||
|
||||
from langchain_typesafe._state import serialize_state
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from langchain_typesafe.types import State
|
||||
|
||||
|
||||
def test_empty_array_and_sequence_are_supported() -> None:
|
||||
"""Empty native arrays and sequences both serialize to an empty array."""
|
||||
assert serialize_state([]) == []
|
||||
assert serialize_state(()) == []
|
||||
|
||||
|
||||
def test_mixed_message_and_json_array_is_serialized_recursively() -> None:
|
||||
"""Messages can appear alongside ordinary JSON values in an array."""
|
||||
state = cast(
|
||||
"State",
|
||||
[HumanMessage("hello"), {"priority": 1}, "plain JSON"],
|
||||
)
|
||||
|
||||
assert serialize_state(state) == [
|
||||
{"role": "user", "content": "hello"},
|
||||
{"priority": 1},
|
||||
"plain JSON",
|
||||
]
|
||||
|
||||
|
||||
def test_messages_can_be_nested_inside_json_objects() -> None:
|
||||
"""Message sequences and individual messages serialize at any object depth."""
|
||||
state: State = {
|
||||
"ticket": {
|
||||
"messages": (
|
||||
SystemMessage("You are a support assistant."),
|
||||
HumanMessage("My integration is broken."),
|
||||
),
|
||||
"draft": AIMessage("I can help troubleshoot it."),
|
||||
},
|
||||
"priority": 2,
|
||||
}
|
||||
|
||||
assert serialize_state(state) == {
|
||||
"ticket": {
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are a support assistant."},
|
||||
{"role": "user", "content": "My integration is broken."},
|
||||
],
|
||||
"draft": {
|
||||
"role": "assistant",
|
||||
"content": "I can help troubleshoot it.",
|
||||
},
|
||||
},
|
||||
"priority": 2,
|
||||
}
|
||||
|
||||
|
||||
def test_unsupported_nested_state_value_is_rejected() -> None:
|
||||
"""Unsupported objects are rejected even when nested in otherwise valid JSON."""
|
||||
state = cast("Any", {"ticket": {"attachment": object()}})
|
||||
|
||||
with pytest.raises(TypeError, match="Unsupported TypeSafe state value: object"):
|
||||
serialize_state(state)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"state",
|
||||
[
|
||||
{1: "value"},
|
||||
{"nested": {1: "value"}},
|
||||
],
|
||||
)
|
||||
def test_non_string_state_key_is_rejected(state: Any) -> None:
|
||||
"""Object keys must remain strings at every state nesting level."""
|
||||
with pytest.raises(TypeError, match="object keys must be strings"):
|
||||
serialize_state(state)
|
||||
|
||||
|
||||
def test_json_tuple_is_serialized_as_array() -> None:
|
||||
"""Python sequences of JSON values become TypeSafe state arrays."""
|
||||
assert serialize_state(("one", 2, None)) == ["one", 2, None]
|
||||
|
||||
|
||||
def test_tool_message_is_serialized_with_tool_context() -> None:
|
||||
"""Tool messages retain their role and tool-call relationship."""
|
||||
state = serialize_state(
|
||||
ToolMessage("Search result", tool_call_id="call_1", name="search")
|
||||
)
|
||||
|
||||
assert state == {
|
||||
"role": "tool",
|
||||
"name": "search",
|
||||
"tool_call_id": "call_1",
|
||||
"content": "Search result",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("state", [None, True, 42, 3.14, b"", bytearray()])
|
||||
def test_invalid_root_state_is_rejected(state: Any) -> None:
|
||||
"""Root state must be a string, object, array, or LangChain message."""
|
||||
with pytest.raises(TypeError, match="TypeSafe state"):
|
||||
serialize_state(state)
|
||||
Generated
+2516
File diff suppressed because it is too large.
Load diff
Reference in new issue
Block a user