mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-17 21:25:46 +02:00
feat(sdk-py): Select-statement (#5933)
This commit is contained in:
@@ -78,6 +78,7 @@ jobs:
|
||||
"libs/checkpoint-sqlite",
|
||||
"libs/checkpoint-postgres",
|
||||
"libs/prebuilt",
|
||||
"libs/sdk-py",
|
||||
]
|
||||
if: needs.changes.outputs.python == 'true' || needs.changes.outputs.deps == 'true'
|
||||
uses: ./.github/workflows/_test.yml
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,7 +1,7 @@
|
||||
.PHONY: lint format
|
||||
.PHONY: lint format test
|
||||
|
||||
test:
|
||||
echo "No tests to run"
|
||||
uv run pytest tests
|
||||
|
||||
######################
|
||||
# LINTING AND FORMATTING
|
||||
|
||||
@@ -33,6 +33,7 @@ import langgraph_sdk
|
||||
from langgraph_sdk.schema import (
|
||||
All,
|
||||
Assistant,
|
||||
AssistantSelectField,
|
||||
AssistantSortBy,
|
||||
AssistantVersion,
|
||||
CancelAction,
|
||||
@@ -41,6 +42,7 @@ from langgraph_sdk.schema import (
|
||||
Config,
|
||||
Context,
|
||||
Cron,
|
||||
CronSelectField,
|
||||
CronSortBy,
|
||||
DisconnectMode,
|
||||
GraphSchema,
|
||||
@@ -54,6 +56,7 @@ from langgraph_sdk.schema import (
|
||||
Run,
|
||||
RunCreate,
|
||||
RunCreateMetadata,
|
||||
RunSelectField,
|
||||
RunStatus,
|
||||
SearchItemsResponse,
|
||||
SortOrder,
|
||||
@@ -61,6 +64,7 @@ from langgraph_sdk.schema import (
|
||||
StreamPart,
|
||||
Subgraphs,
|
||||
Thread,
|
||||
ThreadSelectField,
|
||||
ThreadSortBy,
|
||||
ThreadState,
|
||||
ThreadStatus,
|
||||
@@ -869,6 +873,7 @@ class AssistantsClient:
|
||||
offset: int = 0,
|
||||
sort_by: AssistantSortBy | None = None,
|
||||
sort_order: SortOrder | None = None,
|
||||
select: list[AssistantSelectField] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
) -> list[Assistant]:
|
||||
"""Search for assistants.
|
||||
@@ -910,6 +915,8 @@ class AssistantsClient:
|
||||
payload["sort_by"] = sort_by
|
||||
if sort_order:
|
||||
payload["sort_order"] = sort_order
|
||||
if select:
|
||||
payload["select"] = select
|
||||
return await self.http.post(
|
||||
"/assistants/search",
|
||||
json=payload,
|
||||
@@ -1177,6 +1184,7 @@ class ThreadsClient:
|
||||
offset: int = 0,
|
||||
sort_by: ThreadSortBy | None = None,
|
||||
sort_order: SortOrder | None = None,
|
||||
select: list[ThreadSelectField] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
) -> list[Thread]:
|
||||
"""Search for threads.
|
||||
@@ -1222,6 +1230,8 @@ class ThreadsClient:
|
||||
payload["sort_by"] = sort_by
|
||||
if sort_order:
|
||||
payload["sort_order"] = sort_order
|
||||
if select:
|
||||
payload["select"] = select
|
||||
return await self.http.post(
|
||||
"/threads/search",
|
||||
json=payload,
|
||||
@@ -2146,6 +2156,7 @@ class RunsClient:
|
||||
limit: int = 10,
|
||||
offset: int = 0,
|
||||
status: RunStatus | None = None,
|
||||
select: list[RunSelectField] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
) -> list[Run]:
|
||||
"""List runs.
|
||||
@@ -2178,6 +2189,8 @@ class RunsClient:
|
||||
}
|
||||
if status is not None:
|
||||
params["status"] = status
|
||||
if select:
|
||||
params["select"] = select
|
||||
return await self.http.get(
|
||||
f"/threads/{thread_id}/runs", params=params, headers=headers
|
||||
)
|
||||
@@ -2574,6 +2587,7 @@ class CronClient:
|
||||
offset: int = 0,
|
||||
sort_by: CronSortBy | None = None,
|
||||
sort_order: SortOrder | None = None,
|
||||
select: list[CronSelectField] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
) -> list[Cron]:
|
||||
"""Get a list of cron jobs.
|
||||
@@ -2636,6 +2650,8 @@ class CronClient:
|
||||
payload["sort_by"] = sort_by
|
||||
if sort_order:
|
||||
payload["sort_order"] = sort_order
|
||||
if select:
|
||||
payload["select"] = select
|
||||
payload = {k: v for k, v in payload.items() if v is not None}
|
||||
return await self.http.post("/runs/crons/search", json=payload, headers=headers)
|
||||
|
||||
@@ -3637,6 +3653,7 @@ class SyncAssistantsClient:
|
||||
offset: int = 0,
|
||||
sort_by: AssistantSortBy | None = None,
|
||||
sort_order: SortOrder | None = None,
|
||||
select: list[AssistantSelectField] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
) -> list[Assistant]:
|
||||
"""Search for assistants.
|
||||
@@ -3676,6 +3693,8 @@ class SyncAssistantsClient:
|
||||
payload["sort_by"] = sort_by
|
||||
if sort_order:
|
||||
payload["sort_order"] = sort_order
|
||||
if select:
|
||||
payload["select"] = select
|
||||
return self.http.post(
|
||||
"/assistants/search",
|
||||
json=payload,
|
||||
@@ -3948,6 +3967,7 @@ class SyncThreadsClient:
|
||||
offset: int = 0,
|
||||
sort_by: ThreadSortBy | None = None,
|
||||
sort_order: SortOrder | None = None,
|
||||
select: list[ThreadSelectField] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
) -> list[Thread]:
|
||||
"""Search for threads.
|
||||
@@ -3990,6 +4010,8 @@ class SyncThreadsClient:
|
||||
payload["sort_by"] = sort_by
|
||||
if sort_order:
|
||||
payload["sort_order"] = sort_order
|
||||
if select:
|
||||
payload["select"] = select
|
||||
return self.http.post("/threads/search", json=payload, headers=headers)
|
||||
|
||||
def copy(
|
||||
@@ -4898,6 +4920,8 @@ class SyncRunsClient:
|
||||
*,
|
||||
limit: int = 10,
|
||||
offset: int = 0,
|
||||
status: RunStatus | None = None,
|
||||
select: list[RunSelectField] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
) -> list[Run]:
|
||||
"""List runs.
|
||||
@@ -4923,8 +4947,13 @@ class SyncRunsClient:
|
||||
```
|
||||
|
||||
""" # noqa: E501
|
||||
params: dict[str, Any] = {"limit": limit, "offset": offset}
|
||||
if status is not None:
|
||||
params["status"] = status
|
||||
if select:
|
||||
params["select"] = select
|
||||
return self.http.get(
|
||||
f"/threads/{thread_id}/runs?limit={limit}&offset={offset}", headers=headers
|
||||
f"/threads/{thread_id}/runs", params=params, headers=headers
|
||||
)
|
||||
|
||||
def get(
|
||||
@@ -5319,6 +5348,7 @@ class SyncCronClient:
|
||||
offset: int = 0,
|
||||
sort_by: CronSortBy | None = None,
|
||||
sort_order: SortOrder | None = None,
|
||||
select: list[CronSelectField] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
) -> list[Cron]:
|
||||
"""Get a list of cron jobs.
|
||||
@@ -5380,6 +5410,8 @@ class SyncCronClient:
|
||||
payload["sort_by"] = sort_by
|
||||
if sort_order:
|
||||
payload["sort_order"] = sort_order
|
||||
if select:
|
||||
payload["select"] = select
|
||||
payload = {k: v for k, v in payload.items() if v is not None}
|
||||
return self.http.post("/runs/crons/search", json=payload, headers=headers)
|
||||
|
||||
|
||||
@@ -348,6 +348,62 @@ class Cron(TypedDict):
|
||||
"""The metadata of the cron."""
|
||||
|
||||
|
||||
# Select field aliases for client-side typing of `select` parameters.
|
||||
# These mirror the server's allowed field sets.
|
||||
|
||||
AssistantSelectField = Literal[
|
||||
"assistant_id",
|
||||
"graph_id",
|
||||
"name",
|
||||
"description",
|
||||
"config",
|
||||
"context",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
"metadata",
|
||||
"version",
|
||||
]
|
||||
|
||||
ThreadSelectField = Literal[
|
||||
"thread_id",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
"metadata",
|
||||
"config",
|
||||
"context",
|
||||
"status",
|
||||
"values",
|
||||
"interrupts",
|
||||
]
|
||||
|
||||
RunSelectField = Literal[
|
||||
"run_id",
|
||||
"thread_id",
|
||||
"assistant_id",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
"status",
|
||||
"metadata",
|
||||
"kwargs",
|
||||
"multitask_strategy",
|
||||
]
|
||||
|
||||
CronSelectField = Literal[
|
||||
"cron_id",
|
||||
"assistant_id",
|
||||
"thread_id",
|
||||
"end_time",
|
||||
"schedule",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
"user_id",
|
||||
"payload",
|
||||
"next_run_date",
|
||||
"metadata",
|
||||
"now",
|
||||
]
|
||||
|
||||
|
||||
class RunCreate(TypedDict):
|
||||
"""Defines the parameters for initiating a background run."""
|
||||
|
||||
|
||||
@@ -0,0 +1,76 @@
|
||||
import functools
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import get_args
|
||||
|
||||
from langgraph_sdk.schema import (
|
||||
AssistantSelectField,
|
||||
CronSelectField,
|
||||
RunSelectField,
|
||||
ThreadSelectField,
|
||||
)
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
|
||||
@functools.cache
|
||||
def _load_spec() -> dict:
|
||||
with (
|
||||
Path(current_dir).parents[2]
|
||||
/ "docs"
|
||||
/ "docs"
|
||||
/ "cloud"
|
||||
/ "reference"
|
||||
/ "api"
|
||||
/ "openapi.json"
|
||||
).open() as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def _enum_from_request_select(spec: dict, path: str, method: str) -> set[str]:
|
||||
schema = spec["paths"][path][method]["requestBody"]["content"]["application/json"][
|
||||
"schema"
|
||||
]
|
||||
if "properties" in schema:
|
||||
props = schema["properties"]
|
||||
elif "$ref" in schema:
|
||||
component = spec
|
||||
index = schema["$ref"].split("/")[1:]
|
||||
for part in index:
|
||||
component = component[part]
|
||||
props = component["properties"]
|
||||
else:
|
||||
raise ValueError(f"Unknown schema: {schema}")
|
||||
sel = props["select"]
|
||||
return set(sel["items"]["enum"])
|
||||
|
||||
|
||||
def _enum_from_query_select(spec: dict, path: str, method: str) -> set[str]:
|
||||
params = spec["paths"][path][method]["parameters"]
|
||||
sel = next(p for p in params if p["name"] == "select")
|
||||
return set(sel["schema"]["items"]["enum"])
|
||||
|
||||
|
||||
def test_assistants_select_enum_matches_sdk():
|
||||
spec = _load_spec()
|
||||
expected = set(get_args(AssistantSelectField))
|
||||
assert _enum_from_request_select(spec, "/assistants/search", "post") == expected
|
||||
|
||||
|
||||
def test_threads_select_enum_matches_sdk():
|
||||
spec = _load_spec()
|
||||
expected = set(get_args(ThreadSelectField))
|
||||
assert _enum_from_request_select(spec, "/threads/search", "post") == expected
|
||||
|
||||
|
||||
def test_runs_select_enum_matches_sdk():
|
||||
spec = _load_spec()
|
||||
expected = set(get_args(RunSelectField))
|
||||
assert _enum_from_query_select(spec, "/threads/{thread_id}/runs", "get") == expected
|
||||
|
||||
|
||||
def test_crons_select_enum_matches_sdk():
|
||||
spec = _load_spec()
|
||||
expected = set(get_args(CronSelectField))
|
||||
assert _enum_from_request_select(spec, "/runs/crons/search", "post") == expected
|
||||
Reference in New Issue
Block a user