diff --git a/libs/langgraph-api/package.json b/libs/langgraph-api/package.json index 883bc92..0263925 100644 --- a/libs/langgraph-api/package.json +++ b/libs/langgraph-api/package.json @@ -27,6 +27,10 @@ "types": "./dist/graph/parser/index.d.mts", "default": "./dist/graph/parser/index.mjs" }, + "./embed": { + "types": "./dist/embed.mts", + "default": "./dist/embed.mjs" + }, "./package.json": "./package.json" }, "repository": { diff --git a/libs/langgraph-api/src/api/runs.mts b/libs/langgraph-api/src/api/runs.mts index fbe17e0..7181138 100644 --- a/libs/langgraph-api/src/api/runs.mts +++ b/libs/langgraph-api/src/api/runs.mts @@ -18,16 +18,16 @@ import { serialiseAsDict } from "../utils/serde.mjs"; const api = new Hono(); -const createValidRun = async ( +export const createValidRun = async ( threadId: string | undefined, payload: z.infer, - kwargs: { + kwargs?: { auth: AuthContext | undefined; headers: Headers | undefined; }, ): Promise => { const { assistant_id: assistantId, ...run } = payload; - const { auth, headers } = kwargs; + const { auth, headers } = kwargs ?? {}; const runId = uuid4(); const streamMode = Array.isArray(payload.stream_mode) diff --git a/libs/langgraph-api/src/embed.mts b/libs/langgraph-api/src/embed.mts new file mode 100644 index 0000000..f34a4da --- /dev/null +++ b/libs/langgraph-api/src/embed.mts @@ -0,0 +1,199 @@ +import { BaseCheckpointSaver, BaseStore, Pregel } from "@langchain/langgraph"; +import { Hono } from "hono"; +import { ensureContentType } from "./http/middleware.mjs"; + +import * as schemas from "./schemas.mjs"; + +import { zValidator } from "@hono/zod-validator"; +import { z } from "zod"; +import { streamSSE } from "hono/streaming"; +import type { Metadata, Run } from "./storage/ops.mjs"; +import { streamState } from "./stream.mjs"; +import { serialiseAsDict } from "./utils/serde.mjs"; +import { jsonExtra } from "./utils/hono.mjs"; +import { stateSnapshotToThreadState } from "./state.mjs"; +import { RunnableConfig } from "@langchain/core/runnables"; +import { v4 as uuidv4 } from "uuid"; + +type AnyPregel = Pregel; + +type SimpleThread = { + thread_id: string; + metadata: Metadata; +}; + +export function createServer(app: { + graph: Record; + store?: BaseStore; + checkpointer: BaseCheckpointSaver; + threads: { + get: (threadId: string) => Promise; + put: (threadId: string, options: { metadata?: Metadata }) => Promise; + }; +}) { + const api = new Hono(); + api.use(ensureContentType()); + + api.post("/threads", zValidator("json", schemas.ThreadCreate), async (c) => { + // create a new threaad + const payload = c.req.valid("json"); + const threadId = payload.thread_id || uuidv4(); + + await app.threads.put(threadId, payload); + return jsonExtra(c, { thread_id: threadId }); + }); + + api.get( + "/threads/:thread_id/state", + zValidator("param", z.object({ thread_id: z.string().uuid() })), + zValidator( + "query", + z.object({ subgraphs: schemas.coercedBoolean.optional() }), + ), + async (c) => { + // Get Latest Thread State + const { thread_id } = c.req.valid("param"); + const { subgraphs } = c.req.valid("query"); + + const thread = await app.threads.get(thread_id); + const graphId = thread.metadata?.graph_id as string | undefined | null; + const graph = graphId ? app.graph[graphId] : undefined; + + if (graph == null) { + return jsonExtra( + c, + stateSnapshotToThreadState({ + values: {}, + next: [], + config: {}, + metadata: undefined, + createdAt: undefined, + parentConfig: undefined, + tasks: [], + }), + ); + } + + const config = { configurable: { thread_id } }; + const result = await graph.getState(config, { subgraphs }); + return jsonExtra(c, stateSnapshotToThreadState(result)); + }, + ); + + api.post( + "/threads/:thread_id/history", + zValidator("param", z.object({ thread_id: z.string().uuid() })), + zValidator( + "json", + z.object({ + limit: z.number().optional().default(10), + before: z.string().optional(), + metadata: z.record(z.string(), z.unknown()).optional(), + checkpoint: z + .object({ + checkpoint_id: z.string().uuid().optional(), + checkpoint_ns: z.string().optional(), + checkpoint_map: z.record(z.string(), z.unknown()).optional(), + }) + .optional(), + }), + ), + async (c) => { + const { thread_id } = c.req.valid("param"); + const { limit, before, metadata, checkpoint } = c.req.valid("json"); + + const thread = await app.threads.get(thread_id); + const graphId = thread.metadata?.graph_id as string | undefined | null; + const graph = graphId ? app.graph[graphId] : undefined; + if (graph == null) return jsonExtra(c, []); + + const config = { configurable: { thread_id, ...checkpoint } }; + + const result = []; + const beforeConfig: RunnableConfig | undefined = + typeof before === "string" + ? { configurable: { checkpoint_id: before } } + : before; + + for await (const state of graph.getStateHistory(config, { + limit, + before: beforeConfig, + filter: metadata, + })) { + result.push(stateSnapshotToThreadState(state)); + } + return jsonExtra(c, result); + }, + ); + + api.post( + "/threads/:thread_id/runs/stream", + zValidator("param", z.object({ thread_id: z.string().uuid() })), + zValidator("json", schemas.RunCreate), + async (c) => { + // Stream Run + return streamSSE(c, async (stream) => { + const { thread_id } = c.req.valid("param"); + const payload = c.req.valid("json"); + + const runId = uuidv4(); + const run: Run = { + run_id: runId, + thread_id: thread_id, + assistant_id: payload.assistant_id, + metadata: payload.metadata ?? {}, + status: "running", + kwargs: { + input: payload.input, + command: payload.command, + config: Object.assign( + {}, + payload.config ?? {}, + { + configurable: { + run_id: runId, + thread_id, + graph_id: payload.assistant_id, + }, + }, + { metadata: payload.metadata ?? {} }, + ), + stream_mode: Array.isArray(payload.stream_mode) + ? payload.stream_mode + : payload.stream_mode + ? [payload.stream_mode] + : undefined, + interrupt_before: payload.interrupt_before, + interrupt_after: payload.interrupt_after, + webhook: payload.webhook, + feedback_keys: payload.feedback_keys, + temporary: false, + subgraphs: false, + resumable: false, + }, + multitask_strategy: "reject", + created_at: new Date(), + updated_at: new Date(), + }; + + for await (const { event, data } of streamState(run, 1, { + getGraph: async (graphId) => { + const targetGraph = app.graph[graphId]; + + targetGraph.store = app.store; + targetGraph.checkpointer = app.checkpointer; + + return targetGraph; + }, + })) { + await stream.writeSSE({ + data: serialiseAsDict(data), + event, + }); + } + }); + }, + ); + + return api; +} diff --git a/libs/langgraph-api/src/graph/load.mts b/libs/langgraph-api/src/graph/load.mts index 9b6e43e..bfeb67a 100644 --- a/libs/langgraph-api/src/graph/load.mts +++ b/libs/langgraph-api/src/graph/load.mts @@ -74,6 +74,29 @@ export async function registerFromEnv( ); } +export async function initGraph( + graph: CompiledGraph | CompiledGraphFactory, + config: LangGraphRunnableConfig | undefined, + options?: { + checkpointer?: BaseCheckpointSaver | null; + store?: BaseStore; + }, +) { + const compiled = + typeof graph === "function" + ? await graph(config ?? { configurable: {} }) + : graph; + + if (typeof options?.checkpointer !== "undefined") { + compiled.checkpointer = options?.checkpointer ?? undefined; + } else { + compiled.checkpointer = checkpointer; + } + + compiled.store = options?.store ?? store; + return compiled; +} + export async function getGraph( graphId: string, config: LangGraphRunnableConfig | undefined, @@ -84,21 +107,7 @@ export async function getGraph( ) { if (!GRAPHS[graphId]) throw new HTTPException(404, { message: `Graph "${graphId}" not found` }); - - const compiled = - typeof GRAPHS[graphId] === "function" - ? await GRAPHS[graphId](config ?? { configurable: {} }) - : GRAPHS[graphId]; - - if (typeof options?.checkpointer !== "undefined") { - compiled.checkpointer = options?.checkpointer ?? undefined; - } else { - compiled.checkpointer = checkpointer; - } - - compiled.store = options?.store ?? store; - - return compiled; + return initGraph(GRAPHS[graphId], config, options); } export async function getCachedStaticGraphSchema(graphId: string) { diff --git a/libs/langgraph-api/src/stream.mts b/libs/langgraph-api/src/stream.mts index a7cf1a9..441b3b1 100644 --- a/libs/langgraph-api/src/stream.mts +++ b/libs/langgraph-api/src/stream.mts @@ -1,6 +1,9 @@ import { BaseMessageChunk, isBaseMessage } from "@langchain/core/messages"; import { LangChainTracer } from "@langchain/core/tracers/tracer_langchain"; import { + BaseCheckpointSaver, + BaseStore, + LangGraphRunnableConfig, type CheckpointMetadata, type Interrupt, type StateSnapshot, @@ -147,6 +150,14 @@ export async function* streamState( options?: { onCheckpoint?: (checkpoint: StreamCheckpoint) => void; onTaskResult?: (taskResult: StreamTaskResult) => void; + getGraph?: ( + graphId: string, + config: LangGraphRunnableConfig | undefined, + options?: { + checkpointer?: BaseCheckpointSaver | null; + store?: BaseStore; + }, + ) => Promise>; signal?: AbortSignal; }, ): AsyncGenerator<{ event: string; data: unknown }> { @@ -157,7 +168,7 @@ export async function* streamState( throw new Error("Invalid or missing graph_id"); } - const graph = await getGraph(graphId, kwargs.config, { + const graph = await (options?.getGraph ?? getGraph)(graphId, kwargs.config, { checkpointer: kwargs.temporary ? null : undefined, });