sdk-py: add create_batch method (#988)

This commit is contained in:
Vadym Barda
2024-07-10 17:37:56 -04:00
committed by GitHub
parent dfb2ac321f
commit 8c4da2c41c
+35 -1
View File
@@ -4,7 +4,17 @@ import asyncio
import logging
import os
import sys
from typing import Any, AsyncIterator, Dict, List, NamedTuple, Optional, Union, overload
from typing import (
Any,
AsyncIterator,
Dict,
List,
NamedTuple,
Optional,
TypedDict,
Union,
overload,
)
import httpx
import httpx_sse
@@ -29,6 +39,21 @@ from langgraph_sdk.schema import (
logger = logging.getLogger(__name__)
class RunCreate(TypedDict):
"""Payload for creating a background run."""
thread_id: Optional[str]
assistant_id: str
input: Optional[dict]
metadata: Optional[dict]
config: Optional[Config] = (None,)
checkpoint_id: Optional[str] = (None,)
interrupt_before: Optional[list[str]] = (None,)
interrupt_after: Optional[list[str]] = (None,)
webhook: Optional[str] = (None,)
multitask_strategy: Optional[MultitaskStrategy] = (None,)
def get_client(
*, url: str = "http://localhost:8123", api_key: Optional[str] = None
) -> LangGraphClient:
@@ -561,6 +586,15 @@ class RunsClient:
else:
return await self.http.post("/runs", json=payload)
async def create_batch(self, payloads: list[RunCreate]) -> list[Run]:
"""Create a batch of background runs."""
def filter_payload(payload: RunCreate):
return {k: v for k, v in payload.items() if v is not None}
payloads = [filter_payload(payload) for payload in payloads]
return await self.http.post("/runs/batch", json=payloads)
@overload
async def wait(
self,