ai-infra 0.1.3__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- ai_infra-0.1.3/PKG-INFO +39 -0
- ai_infra-0.1.3/README.md +1 -0
- ai_infra-0.1.3/pyproject.toml +59 -0
- ai_infra-0.1.3/src/ai_infra/__init__.py +20 -0
- ai_infra-0.1.3/src/ai_infra/graph/__init__.py +8 -0
- ai_infra-0.1.3/src/ai_infra/graph/core.py +156 -0
- ai_infra-0.1.3/src/ai_infra/graph/examples/01_graph_basic.py +45 -0
- ai_infra-0.1.3/src/ai_infra/graph/examples/02_graph_stream_values.py +27 -0
- ai_infra-0.1.3/src/ai_infra/graph/examples/__init__.py +0 -0
- ai_infra-0.1.3/src/ai_infra/graph/models.py +35 -0
- ai_infra-0.1.3/src/ai_infra/graph/utils.py +139 -0
- ai_infra-0.1.3/src/ai_infra/llm/__init__.py +13 -0
- ai_infra-0.1.3/src/ai_infra/llm/core.py +357 -0
- ai_infra-0.1.3/src/ai_infra/llm/examples/01_agent_basic.py +16 -0
- ai_infra-0.1.3/src/ai_infra/llm/examples/02_llm_chat_basic.py +16 -0
- ai_infra-0.1.3/src/ai_infra/llm/examples/03_structured_output.py +27 -0
- ai_infra-0.1.3/src/ai_infra/llm/examples/04_agent_stream.py +29 -0
- ai_infra-0.1.3/src/ai_infra/llm/examples/05_tool_controls.py +34 -0
- ai_infra-0.1.3/src/ai_infra/llm/examples/06_hitl.py +43 -0
- ai_infra-0.1.3/src/ai_infra/llm/examples/07_retry.py +27 -0
- ai_infra-0.1.3/src/ai_infra/llm/examples/08_agent_stream_tokens.py +22 -0
- ai_infra-0.1.3/src/ai_infra/llm/examples/09_chat_stream.py +25 -0
- ai_infra-0.1.3/src/ai_infra/llm/examples/__init__.py +0 -0
- ai_infra-0.1.3/src/ai_infra/llm/providers/__init__.py +2 -0
- ai_infra-0.1.3/src/ai_infra/llm/providers/models.py +28 -0
- ai_infra-0.1.3/src/ai_infra/llm/providers/providers.py +5 -0
- ai_infra-0.1.3/src/ai_infra/llm/tools/__init__.py +2 -0
- ai_infra-0.1.3/src/ai_infra/llm/tools/tool_controls.py +114 -0
- ai_infra-0.1.3/src/ai_infra/llm/tools/tools.py +215 -0
- ai_infra-0.1.3/src/ai_infra/llm/utils/__init__.py +38 -0
- ai_infra-0.1.3/src/ai_infra/llm/utils/fallbacks.py +115 -0
- ai_infra-0.1.3/src/ai_infra/llm/utils/messages.py +24 -0
- ai_infra-0.1.3/src/ai_infra/llm/utils/model_init.py +24 -0
- ai_infra-0.1.3/src/ai_infra/llm/utils/retry.py +16 -0
- ai_infra-0.1.3/src/ai_infra/llm/utils/runtime_bind.py +169 -0
- ai_infra-0.1.3/src/ai_infra/llm/utils/settings.py +9 -0
- ai_infra-0.1.3/src/ai_infra/llm/utils/validation.py +23 -0
- ai_infra-0.1.3/src/ai_infra/mcp/__init__.py +7 -0
- ai_infra-0.1.3/src/ai_infra/mcp/client/__init__.py +0 -0
- ai_infra-0.1.3/src/ai_infra/mcp/client/core.py +460 -0
- ai_infra-0.1.3/src/ai_infra/mcp/client/models.py +29 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/01_mcps.py +9 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/__init__.py +0 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/agents/01_streamable_http_agent.py +28 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/agents/02_multi_server_agent.py +32 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/agents/__init__.py +0 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/client/01_sse.py +18 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/client/02_stdio.py +20 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/client/03_streamable_http.py +22 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/client/04_stdio.py +23 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/client/05_openapi.py +15 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/client/06_multi_server_client.py +20 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/client/06_server_metadata_from_client.py +30 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/client/__init__.py +0 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/resources/apiframeworks.json +3 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/resources/spotify.yaml +6954 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/server/01_sse.py +10 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/server/02_stdio.py +12 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/server/03_streamable_http.py +11 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/server/04_openapi.py +29 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/server/__init__.py +0 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/server/fastapi/01_add_app.py +38 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/server/fastapi/02_add_fastmcp.py +32 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/server/fastapi/03_raw_mount.py +29 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/server/fastapi/__init__.py +0 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/server/fastapi/from_module/01_from_module_fastmcp.py +9 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/server/fastapi/from_module/02_from_modle_asgi.py +13 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/server/fastapi/from_module/03_from_module.py +22 -0
- ai_infra-0.1.3/src/ai_infra/mcp/examples/server/fastapi/from_module/__init__.py +0 -0
- ai_infra-0.1.3/src/ai_infra/mcp/server/__init__.py +0 -0
- ai_infra-0.1.3/src/ai_infra/mcp/server/core.py +379 -0
- ai_infra-0.1.3/src/ai_infra/mcp/server/models.py +12 -0
- ai_infra-0.1.3/src/ai_infra/mcp/server/openapi/__init__.py +4 -0
- ai_infra-0.1.3/src/ai_infra/mcp/server/openapi/builder.py +582 -0
- ai_infra-0.1.3/src/ai_infra/mcp/server/openapi/constants.py +3 -0
- ai_infra-0.1.3/src/ai_infra/mcp/server/openapi/io.py +19 -0
- ai_infra-0.1.3/src/ai_infra/mcp/server/openapi/models.py +60 -0
- ai_infra-0.1.3/src/ai_infra/mcp/server/openapi/runtime.py +90 -0
- ai_infra-0.1.3/src/ai_infra/mcp/server/tools.py +56 -0
- ai_infra-0.1.3/src/ai_infra/py.typed +0 -0
ai_infra-0.1.3/PKG-INFO
ADDED
|
@@ -0,0 +1,39 @@
|
|
|
1
|
+
Metadata-Version: 2.3
|
|
2
|
+
Name: ai-infra
|
|
3
|
+
Version: 0.1.3
|
|
4
|
+
Summary: Infrastructure for efficient and scalable AI applications.
|
|
5
|
+
License: MIT
|
|
6
|
+
Keywords: ai,langchain,langgraph,fastapi,infra,llm,mcp
|
|
7
|
+
Author: Ali Khatami
|
|
8
|
+
Author-email: aliikhatami94@gmail.com
|
|
9
|
+
Requires-Python: >=3.11,<4.0
|
|
10
|
+
Classifier: Development Status :: 4 - Beta
|
|
11
|
+
Classifier: Framework :: FastAPI
|
|
12
|
+
Classifier: Intended Audience :: Developers
|
|
13
|
+
Classifier: License :: OSI Approved :: MIT License
|
|
14
|
+
Classifier: Programming Language :: Python :: 3
|
|
15
|
+
Classifier: Programming Language :: Python :: 3.11
|
|
16
|
+
Classifier: Programming Language :: Python :: 3.12
|
|
17
|
+
Classifier: Programming Language :: Python :: 3.13
|
|
18
|
+
Classifier: Programming Language :: Python :: 3 :: Only
|
|
19
|
+
Classifier: Typing :: Typed
|
|
20
|
+
Requires-Dist: fastapi (>=0.116.1,<0.117.0)
|
|
21
|
+
Requires-Dist: langchain (>=0.3.27,<0.4.0)
|
|
22
|
+
Requires-Dist: langchain-anthropic (>=0.3.18,<0.4.0)
|
|
23
|
+
Requires-Dist: langchain-google-genai (>=2.1.9,<3.0.0)
|
|
24
|
+
Requires-Dist: langchain-mcp-adapters (>=0.1.9,<0.2.0)
|
|
25
|
+
Requires-Dist: langchain-openai (>=0.3.29,<0.4.0)
|
|
26
|
+
Requires-Dist: langchain-xai (>=0.2.5,<0.3.0)
|
|
27
|
+
Requires-Dist: langgraph (>=0.6.4,<0.7.0)
|
|
28
|
+
Requires-Dist: langsmith (>=0.4.13,<0.5.0)
|
|
29
|
+
Requires-Dist: mcp[cli] (>=1.13.1,<2.0.0)
|
|
30
|
+
Requires-Dist: pydantic (>=2.11,<3.0)
|
|
31
|
+
Requires-Dist: python-dotenv (>=1.1.1,<2.0.0)
|
|
32
|
+
Project-URL: Documentation, https://github.com/your-org/ai-infra#readme
|
|
33
|
+
Project-URL: Homepage, https://github.com/your-org/ai-infra
|
|
34
|
+
Project-URL: Issues, https://github.com/your-org/ai-infra/issues
|
|
35
|
+
Project-URL: Repository, https://github.com/your-org/ai-infra
|
|
36
|
+
Description-Content-Type: text/markdown
|
|
37
|
+
|
|
38
|
+
# ai-infra
|
|
39
|
+
|
ai_infra-0.1.3/README.md
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
# ai-infra
|
|
@@ -0,0 +1,59 @@
|
|
|
1
|
+
[tool.poetry]
|
|
2
|
+
name = "ai-infra"
|
|
3
|
+
version = "0.1.3"
|
|
4
|
+
description = "Infrastructure for efficient and scalable AI applications."
|
|
5
|
+
authors = ["Ali Khatami <aliikhatami94@gmail.com>"]
|
|
6
|
+
license = "MIT"
|
|
7
|
+
readme = "README.md"
|
|
8
|
+
packages = [{ include = "ai_infra", from = "src" }]
|
|
9
|
+
keywords = ["ai", "langchain", "langgraph", "fastapi", "infra", "llm", "mcp"]
|
|
10
|
+
|
|
11
|
+
classifiers = [
|
|
12
|
+
"Development Status :: 4 - Beta",
|
|
13
|
+
"Intended Audience :: Developers",
|
|
14
|
+
"License :: OSI Approved :: MIT License",
|
|
15
|
+
"Programming Language :: Python :: 3",
|
|
16
|
+
"Programming Language :: Python :: 3 :: Only",
|
|
17
|
+
"Programming Language :: Python :: 3.11",
|
|
18
|
+
"Programming Language :: Python :: 3.12",
|
|
19
|
+
"Programming Language :: Python :: 3.13",
|
|
20
|
+
"Framework :: FastAPI",
|
|
21
|
+
"Typing :: Typed"
|
|
22
|
+
]
|
|
23
|
+
|
|
24
|
+
[tool.poetry.urls]
|
|
25
|
+
Homepage = "https://github.com/your-org/ai-infra"
|
|
26
|
+
Repository = "https://github.com/your-org/ai-infra"
|
|
27
|
+
Issues = "https://github.com/your-org/ai-infra/issues"
|
|
28
|
+
Documentation = "https://github.com/your-org/ai-infra#readme"
|
|
29
|
+
|
|
30
|
+
[tool.poetry.dependencies]
|
|
31
|
+
python = ">=3.11,<4.0"
|
|
32
|
+
|
|
33
|
+
# Core AI/Infra
|
|
34
|
+
fastapi = "^0.116.1"
|
|
35
|
+
pydantic = "^2.11"
|
|
36
|
+
python-dotenv = "^1.1.1"
|
|
37
|
+
|
|
38
|
+
# LangChain ecosystem
|
|
39
|
+
langgraph = ">=0.6.4,<0.7.0"
|
|
40
|
+
langchain = "^0.3.27"
|
|
41
|
+
langchain-openai = "^0.3.29"
|
|
42
|
+
langchain-xai = "^0.2.5"
|
|
43
|
+
langchain-anthropic = "^0.3.18"
|
|
44
|
+
langchain-google-genai = "^2.1.9"
|
|
45
|
+
langchain-mcp-adapters = "^0.1.9"
|
|
46
|
+
langsmith = "^0.4.13"
|
|
47
|
+
|
|
48
|
+
# MCP integration
|
|
49
|
+
mcp = {extras = ["cli"], version = "^1.13.1"}
|
|
50
|
+
|
|
51
|
+
[tool.poetry.group.dev.dependencies]
|
|
52
|
+
pytest = "^8.3.0"
|
|
53
|
+
pytest-asyncio = "^0.23.0"
|
|
54
|
+
ruff = "^0.5.0"
|
|
55
|
+
mypy = "^1.10.0"
|
|
56
|
+
|
|
57
|
+
[build-system]
|
|
58
|
+
requires = ["poetry-core>=2.0.0,<3.0.0"]
|
|
59
|
+
build-backend = "poetry.core.masonry.api"
|
|
@@ -0,0 +1,20 @@
|
|
|
1
|
+
import os
|
|
2
|
+
from dotenv import load_dotenv, find_dotenv
|
|
3
|
+
|
|
4
|
+
if not os.environ.get("AI_INFRA_ENV_LOADED"):
|
|
5
|
+
load_dotenv(find_dotenv(usecwd=True))
|
|
6
|
+
os.environ["AI_INFRA_ENV_LOADED"] = "1"
|
|
7
|
+
|
|
8
|
+
# Re-export primary public API components
|
|
9
|
+
from ai_infra.llm.core import CoreLLM
|
|
10
|
+
from ai_infra.graph.core import CoreGraph
|
|
11
|
+
from ai_infra.llm.providers import Providers
|
|
12
|
+
from ai_infra.llm.providers.models import Models
|
|
13
|
+
|
|
14
|
+
__all__ = [
|
|
15
|
+
"CoreGraph",
|
|
16
|
+
"Models",
|
|
17
|
+
"Providers",
|
|
18
|
+
"CoreMCP",
|
|
19
|
+
]
|
|
20
|
+
|
|
@@ -0,0 +1,156 @@
|
|
|
1
|
+
from collections.abc import AsyncIterator, Iterator
|
|
2
|
+
from typing import Any, Sequence, Union, Dict
|
|
3
|
+
from langgraph.constants import START, END
|
|
4
|
+
from langgraph.graph import StateGraph
|
|
5
|
+
|
|
6
|
+
from ai_infra.graph.models import GraphStructure, EdgeType
|
|
7
|
+
from ai_infra.graph.utils import (
|
|
8
|
+
normalize_node_definitions, normalize_initial_state,
|
|
9
|
+
build_edges, wrap_node, normalize_stream_mode,
|
|
10
|
+
make_hook, make_trace_fn, make_trace_wrapper
|
|
11
|
+
)
|
|
12
|
+
|
|
13
|
+
class CoreGraph:
|
|
14
|
+
def __init__(
|
|
15
|
+
self,
|
|
16
|
+
*,
|
|
17
|
+
state_type: type,
|
|
18
|
+
node_definitions: Union[Sequence, dict],
|
|
19
|
+
edges: Sequence[EdgeType],
|
|
20
|
+
checkpointer=None,
|
|
21
|
+
store=None
|
|
22
|
+
):
|
|
23
|
+
if not (isinstance(state_type, type) and (issubclass(state_type, dict) or hasattr(state_type, '__annotations__'))):
|
|
24
|
+
raise ValueError("state_type must be a TypedDict or dict subclass")
|
|
25
|
+
self.state_type = state_type
|
|
26
|
+
|
|
27
|
+
node_definitions = normalize_node_definitions(node_definitions)
|
|
28
|
+
self.node_definitions = list(node_definitions.items())
|
|
29
|
+
|
|
30
|
+
# centralize edge building/validation + START/END
|
|
31
|
+
regular_edges, conditional_edges = build_edges(list(node_definitions.keys()), edges)
|
|
32
|
+
self.edges = regular_edges
|
|
33
|
+
self.conditional_edges = conditional_edges
|
|
34
|
+
|
|
35
|
+
self._checkpointer = checkpointer
|
|
36
|
+
self._store = store
|
|
37
|
+
self.graph = self._build_graph().compile(checkpointer=self._checkpointer, store=self._store)
|
|
38
|
+
|
|
39
|
+
def _build_graph(self, node_items=None, sync: bool=False) -> StateGraph:
|
|
40
|
+
wf = StateGraph(self.state_type)
|
|
41
|
+
node_items = node_items or self.node_definitions
|
|
42
|
+
for name, fn in node_items:
|
|
43
|
+
wf.add_node(name, wrap_node(fn, sync))
|
|
44
|
+
for start, router_fn, path_map in self.conditional_edges:
|
|
45
|
+
wf.add_conditional_edges(start, wrap_node(router_fn, sync), path_map)
|
|
46
|
+
for start, end in self.edges:
|
|
47
|
+
wf.add_edge(start, end)
|
|
48
|
+
return wf
|
|
49
|
+
|
|
50
|
+
def _prepare_run(
|
|
51
|
+
self,
|
|
52
|
+
initial_state=None,
|
|
53
|
+
*,
|
|
54
|
+
config=None,
|
|
55
|
+
on_enter=None,
|
|
56
|
+
on_exit=None,
|
|
57
|
+
trace=None,
|
|
58
|
+
sync: bool=False,
|
|
59
|
+
**kwargs
|
|
60
|
+
):
|
|
61
|
+
initial_state = normalize_initial_state(initial_state, kwargs)
|
|
62
|
+
# fast path: no hooks → use cached compiled graph
|
|
63
|
+
if not (on_enter or on_exit or trace):
|
|
64
|
+
compiled = self._build_graph(sync=sync).compile(checkpointer=self._checkpointer, store=self._store) if sync else self.graph
|
|
65
|
+
return compiled, initial_state, config
|
|
66
|
+
|
|
67
|
+
# else, patch & recompile
|
|
68
|
+
on_enter_fn = make_hook(on_enter, sync=sync)
|
|
69
|
+
on_exit_fn = make_hook(on_exit, sync=sync)
|
|
70
|
+
trace_fn = make_trace_fn(trace, sync=sync)
|
|
71
|
+
patched_nodes = [
|
|
72
|
+
(name, make_trace_wrapper(name, wrap_node(fn, sync), on_enter_fn, on_exit_fn, trace_fn, sync))
|
|
73
|
+
for name, fn in self.node_definitions
|
|
74
|
+
]
|
|
75
|
+
compiled = self._build_graph(node_items=patched_nodes, sync=sync).compile(checkpointer=self._checkpointer, store=self._store)
|
|
76
|
+
return compiled, initial_state, config
|
|
77
|
+
|
|
78
|
+
# ---- invoke -----------------------------------------------------------------
|
|
79
|
+
async def arun(self, initial_state=None, *, config=None, on_enter=None, on_exit=None, trace=None, **kwargs) -> Any:
|
|
80
|
+
compiled, initial_state, config = self._prepare_run(initial_state, config=config, on_enter=on_enter, on_exit=on_exit, trace=trace, sync=False, **kwargs)
|
|
81
|
+
return await compiled.ainvoke(initial_state, config=config) if config is not None else await compiled.ainvoke(initial_state)
|
|
82
|
+
|
|
83
|
+
def run(self, initial_state=None, *, config=None, on_enter=None, on_exit=None, trace=None, **kwargs) -> Any:
|
|
84
|
+
compiled, initial_state, config = self._prepare_run(initial_state, config=config, on_enter=on_enter, on_exit=on_exit, trace=trace, sync=True, **kwargs)
|
|
85
|
+
return compiled.invoke(initial_state, config=config) if config is not None else compiled.invoke(initial_state)
|
|
86
|
+
|
|
87
|
+
# ---- streaming ---------------------------------------------------------------
|
|
88
|
+
async def astream(self, initial_state=None, *, config=None, stream_mode=("updates","values")) -> AsyncIterator[tuple[str, Any]]:
|
|
89
|
+
stream_mode = normalize_stream_mode(stream_mode)
|
|
90
|
+
compiled, initial_state, config = self._prepare_run(initial_state, config=config, sync=False)
|
|
91
|
+
async for mode, chunk in compiled.astream(initial_state, config=config, stream_mode=stream_mode):
|
|
92
|
+
yield mode, chunk
|
|
93
|
+
|
|
94
|
+
def stream(self, initial_state=None, *, config=None, stream_mode=("updates","values")) -> Iterator[tuple[str, Any]]:
|
|
95
|
+
stream_mode = normalize_stream_mode(stream_mode)
|
|
96
|
+
compiled, initial_state, config = self._prepare_run(initial_state, config=config, sync=True)
|
|
97
|
+
for mode, chunk in compiled.stream(initial_state, config=config, stream_mode=stream_mode):
|
|
98
|
+
yield mode, chunk
|
|
99
|
+
|
|
100
|
+
async def astream_values(self, initial_state=None, *, config=None):
|
|
101
|
+
async for _, chunk in self.astream(initial_state, config=config, stream_mode="values"):
|
|
102
|
+
yield chunk
|
|
103
|
+
|
|
104
|
+
def stream_values(self, initial_state=None, *, config=None):
|
|
105
|
+
for _, chunk in self.stream(initial_state, config=config, stream_mode="values"):
|
|
106
|
+
yield chunk
|
|
107
|
+
|
|
108
|
+
# ---- analysis / debug --------------------------------------------------------
|
|
109
|
+
def analyze(self) -> GraphStructure:
|
|
110
|
+
nodes = [name for name, _ in self.node_definitions]
|
|
111
|
+
entry_points = [end for start, end in self.edges if start == START] or nodes[:1]
|
|
112
|
+
exit_points = [start for start, end in self.edges if end == END]
|
|
113
|
+
conditional_edges_data = [
|
|
114
|
+
{"start": start, "router_function": getattr(router_fn, '__name__', str(router_fn)), "path_options": list(path_map.keys())}
|
|
115
|
+
for start, router_fn, path_map in self.conditional_edges
|
|
116
|
+
] if self.conditional_edges else None
|
|
117
|
+
state_schema = {key: getattr(value, '__name__', str(value)) for key, value in getattr(self.state_type, '__annotations__', {}).items()}
|
|
118
|
+
# reachability
|
|
119
|
+
reachable = set(entry_points)
|
|
120
|
+
edges_map = {start: [] for start, _ in self.edges}
|
|
121
|
+
for start, end in self.edges:
|
|
122
|
+
edges_map.setdefault(start, []).append(end)
|
|
123
|
+
queue = list(entry_points)
|
|
124
|
+
while queue:
|
|
125
|
+
node = queue.pop(0)
|
|
126
|
+
for nbr in edges_map.get(node, []):
|
|
127
|
+
if nbr not in reachable and nbr not in (START, END):
|
|
128
|
+
reachable.add(nbr); queue.append(nbr)
|
|
129
|
+
unreachable = [n for n in nodes if n not in reachable]
|
|
130
|
+
|
|
131
|
+
return GraphStructure(
|
|
132
|
+
state_type_name=self.state_type.__name__,
|
|
133
|
+
state_schema=state_schema,
|
|
134
|
+
node_count=len(nodes),
|
|
135
|
+
nodes=nodes,
|
|
136
|
+
edge_count=len(self.edges),
|
|
137
|
+
edges=self.edges,
|
|
138
|
+
conditional_edge_count=len(self.conditional_edges) if self.conditional_edges else 0,
|
|
139
|
+
conditional_edges=conditional_edges_data,
|
|
140
|
+
entry_points=entry_points,
|
|
141
|
+
exit_points=exit_points,
|
|
142
|
+
has_memory=self._checkpointer is not None,
|
|
143
|
+
unreachable=unreachable
|
|
144
|
+
)
|
|
145
|
+
|
|
146
|
+
def describe(self) -> Dict:
|
|
147
|
+
return self.analyze().model_dump()
|
|
148
|
+
|
|
149
|
+
def get_state(self, config):
|
|
150
|
+
return self.graph.get_state(config)
|
|
151
|
+
|
|
152
|
+
def get_state_history(self, config):
|
|
153
|
+
return list(self.graph.get_state_history(config))
|
|
154
|
+
|
|
155
|
+
def get_arch_diagram(self) -> str:
|
|
156
|
+
return self.graph.get_graph().draw_mermaid()
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
"""01_graph_basic: Basic graph with conditional looping.
|
|
2
|
+
Usage: python -m quickstart.run graph_basic
|
|
3
|
+
"""
|
|
4
|
+
from typing_extensions import TypedDict
|
|
5
|
+
from langgraph.graph import END
|
|
6
|
+
from ai_infra.graph.core import CoreGraph
|
|
7
|
+
from ai_infra.graph.models import Edge, ConditionalEdge
|
|
8
|
+
|
|
9
|
+
MAX_VALUE = 40
|
|
10
|
+
|
|
11
|
+
class MyState(TypedDict):
|
|
12
|
+
value: int
|
|
13
|
+
|
|
14
|
+
def inc(state: MyState) -> MyState:
|
|
15
|
+
"""Increment value."""
|
|
16
|
+
state["value"] += 1
|
|
17
|
+
return state
|
|
18
|
+
|
|
19
|
+
def mul(state: MyState) -> MyState:
|
|
20
|
+
"""Double value."""
|
|
21
|
+
state["value"] *= 2
|
|
22
|
+
return state
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def _trace(node_name, state, event): # type: ignore[override]
|
|
26
|
+
print(f"{event.upper()} node={node_name} state={state}")
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
graph = CoreGraph(
|
|
30
|
+
state_type=MyState,
|
|
31
|
+
node_definitions=[inc, mul],
|
|
32
|
+
edges=[
|
|
33
|
+
Edge(start="inc", end="mul"),
|
|
34
|
+
ConditionalEdge(
|
|
35
|
+
start="mul",
|
|
36
|
+
router_fn=lambda s: "inc" if s["value"] < MAX_VALUE else END,
|
|
37
|
+
targets=["inc", END],
|
|
38
|
+
),
|
|
39
|
+
],
|
|
40
|
+
)
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def main():
|
|
44
|
+
result = graph.run({"value": 1}, trace=_trace)
|
|
45
|
+
print("Final:", result)
|
|
@@ -0,0 +1,27 @@
|
|
|
1
|
+
"""02_graph_stream_values: Stream only state value snapshots.
|
|
2
|
+
Usage: python -m quickstart.run graph_stream_values
|
|
3
|
+
"""
|
|
4
|
+
from typing_extensions import TypedDict
|
|
5
|
+
from ai_infra.graph.core import CoreGraph
|
|
6
|
+
from ai_infra.graph.models import Edge
|
|
7
|
+
|
|
8
|
+
class MyState(TypedDict):
|
|
9
|
+
value: int
|
|
10
|
+
|
|
11
|
+
def inc(state: MyState) -> MyState:
|
|
12
|
+
state["value"] += 1
|
|
13
|
+
return state
|
|
14
|
+
|
|
15
|
+
def main():
|
|
16
|
+
graph = CoreGraph(
|
|
17
|
+
state_type=MyState,
|
|
18
|
+
node_definitions=[inc],
|
|
19
|
+
edges=[Edge(start="inc", end="inc")], # simple loop; rely on user to break (example)
|
|
20
|
+
)
|
|
21
|
+
# For demonstration, manually break after 5 iterations
|
|
22
|
+
iterations = 0
|
|
23
|
+
for snapshot in graph.stream_values({"value": 0}):
|
|
24
|
+
print(snapshot)
|
|
25
|
+
iterations += 1
|
|
26
|
+
if iterations >= 5:
|
|
27
|
+
break
|
|
File without changes
|
|
@@ -0,0 +1,35 @@
|
|
|
1
|
+
from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple, Union
|
|
2
|
+
from pydantic import BaseModel, ConfigDict
|
|
3
|
+
|
|
4
|
+
class GraphStructure(BaseModel):
|
|
5
|
+
"""Pydantic model representing the graph's structural information."""
|
|
6
|
+
state_type_name: str
|
|
7
|
+
state_schema: Dict[str, str]
|
|
8
|
+
node_count: int
|
|
9
|
+
nodes: List[str]
|
|
10
|
+
edge_count: int
|
|
11
|
+
edges: List[Tuple[str, str]]
|
|
12
|
+
conditional_edge_count: int
|
|
13
|
+
conditional_edges: Optional[List[Dict[str, Any]]] = None
|
|
14
|
+
entry_points: List[str]
|
|
15
|
+
exit_points: List[str]
|
|
16
|
+
has_memory: bool
|
|
17
|
+
unreachable: Optional[List[str]] = None
|
|
18
|
+
|
|
19
|
+
class CoreGraphConfig(BaseModel):
|
|
20
|
+
model_config = ConfigDict(arbitrary_types_allowed=True)
|
|
21
|
+
node_definitions: Sequence[Any]
|
|
22
|
+
edges: Sequence[Tuple[str, str]]
|
|
23
|
+
conditional_edges: Optional[Sequence[Tuple[str, Any, dict]]] = None
|
|
24
|
+
memory_store: Optional[object] = None
|
|
25
|
+
|
|
26
|
+
class Edge(BaseModel):
|
|
27
|
+
start: str
|
|
28
|
+
end: str
|
|
29
|
+
|
|
30
|
+
class ConditionalEdge(BaseModel):
|
|
31
|
+
start: str
|
|
32
|
+
router_fn: Callable
|
|
33
|
+
targets: list[str]
|
|
34
|
+
|
|
35
|
+
EdgeType = Union[Edge, ConditionalEdge]
|
|
@@ -0,0 +1,139 @@
|
|
|
1
|
+
import inspect
|
|
2
|
+
import asyncio
|
|
3
|
+
from typing import Sequence, Any
|
|
4
|
+
from langgraph.constants import START, END
|
|
5
|
+
from ai_infra.graph.models import Edge, ConditionalEdge
|
|
6
|
+
|
|
7
|
+
def normalize_node_definitions(node_definitions):
|
|
8
|
+
if isinstance(node_definitions, dict):
|
|
9
|
+
return node_definitions.copy()
|
|
10
|
+
return {fn.__name__: fn for fn in node_definitions}
|
|
11
|
+
|
|
12
|
+
def normalize_initial_state(initial_state, kwargs):
|
|
13
|
+
if initial_state is None:
|
|
14
|
+
return kwargs
|
|
15
|
+
if kwargs:
|
|
16
|
+
raise ValueError("Provide either initial_state or keyword arguments, not both.")
|
|
17
|
+
return initial_state
|
|
18
|
+
|
|
19
|
+
def validate_edges(edges, all_nodes):
|
|
20
|
+
for start, end in edges:
|
|
21
|
+
for endpoint in (start, end):
|
|
22
|
+
if endpoint not in all_nodes and endpoint not in (START, END):
|
|
23
|
+
raise ValueError(f"Edge endpoint '{endpoint}' is not a known node or START/END")
|
|
24
|
+
|
|
25
|
+
def validate_conditional_edges(conditional_edges, all_nodes):
|
|
26
|
+
for start, router_fn, path_map in conditional_edges:
|
|
27
|
+
if start not in all_nodes and start not in (START, END):
|
|
28
|
+
raise ValueError(f"Conditional edge start '{start}' is not a known node or START/END")
|
|
29
|
+
for target in path_map.values():
|
|
30
|
+
if target not in all_nodes and target not in (START, END):
|
|
31
|
+
raise ValueError(f"Conditional path target '{target}' is not a known node or START/END")
|
|
32
|
+
|
|
33
|
+
def make_router_wrapper(fn, valid_targets):
|
|
34
|
+
async def wrapper(state):
|
|
35
|
+
result = await fn(state) if inspect.iscoroutinefunction(fn) else fn(state)
|
|
36
|
+
if result not in valid_targets:
|
|
37
|
+
raise ValueError(f"Router function returned '{result}', which is not in targets {valid_targets}")
|
|
38
|
+
return result
|
|
39
|
+
return wrapper
|
|
40
|
+
|
|
41
|
+
def make_hook(hook, event=None, sync=False):
|
|
42
|
+
if not hook:
|
|
43
|
+
return None
|
|
44
|
+
if inspect.iscoroutinefunction(hook):
|
|
45
|
+
if sync:
|
|
46
|
+
def sync_hook(node, state):
|
|
47
|
+
return asyncio.run(hook(node, state) if event is None else hook(node, state, event))
|
|
48
|
+
return sync_hook
|
|
49
|
+
return lambda node, state: hook(node, state) if event is None else hook(node, state, event)
|
|
50
|
+
async def async_hook(node, state):
|
|
51
|
+
return hook(node, state) if event is None else hook(node, state, event)
|
|
52
|
+
return async_hook
|
|
53
|
+
|
|
54
|
+
def make_trace_fn(trace, sync=False):
|
|
55
|
+
if not trace:
|
|
56
|
+
return None
|
|
57
|
+
if sync:
|
|
58
|
+
def trace_sync(node, state, event):
|
|
59
|
+
return asyncio.run(trace(node, state, event)) if inspect.iscoroutinefunction(trace) else trace(node, state, event)
|
|
60
|
+
return trace_sync
|
|
61
|
+
async def trace_async(node, state, event):
|
|
62
|
+
if inspect.iscoroutinefunction(trace):
|
|
63
|
+
await trace(node, state, event)
|
|
64
|
+
else:
|
|
65
|
+
trace(node, state, event)
|
|
66
|
+
return trace_async
|
|
67
|
+
|
|
68
|
+
def make_trace_wrapper(name, fn, on_enter, on_exit, trace, sync):
|
|
69
|
+
if sync:
|
|
70
|
+
def wrapped(state):
|
|
71
|
+
if on_enter: on_enter(name, state)
|
|
72
|
+
if trace: trace(name, state, "enter")
|
|
73
|
+
result = fn(state)
|
|
74
|
+
if on_exit: on_exit(name, result)
|
|
75
|
+
if trace: trace(name, result, "exit")
|
|
76
|
+
return result
|
|
77
|
+
return wrapped
|
|
78
|
+
async def wrapped(state):
|
|
79
|
+
if on_enter: await on_enter(name, state)
|
|
80
|
+
if trace: await trace(name, state, "enter")
|
|
81
|
+
result = await fn(state)
|
|
82
|
+
if on_exit: await on_exit(name, result)
|
|
83
|
+
if trace: await trace(name, result, "exit")
|
|
84
|
+
return result
|
|
85
|
+
return wrapped
|
|
86
|
+
|
|
87
|
+
# ---- new helpers ---------------------------------------------------------------
|
|
88
|
+
|
|
89
|
+
def normalize_stream_mode(stream_mode):
|
|
90
|
+
if stream_mode is None:
|
|
91
|
+
return ["updates"]
|
|
92
|
+
if isinstance(stream_mode, str):
|
|
93
|
+
return [stream_mode]
|
|
94
|
+
return list(stream_mode)
|
|
95
|
+
|
|
96
|
+
def wrap_node(fn, sync: bool):
|
|
97
|
+
if sync:
|
|
98
|
+
if not inspect.iscoroutinefunction(fn):
|
|
99
|
+
return fn
|
|
100
|
+
def sync_wrapper(*args, **kwargs):
|
|
101
|
+
try:
|
|
102
|
+
loop = asyncio.get_running_loop()
|
|
103
|
+
if loop.is_running():
|
|
104
|
+
raise RuntimeError(
|
|
105
|
+
"CoreGraph.run/stream cannot execute async nodes inside a running event loop. "
|
|
106
|
+
"Use arun/astream instead."
|
|
107
|
+
)
|
|
108
|
+
except RuntimeError:
|
|
109
|
+
# no running loop; safe to asyncio.run
|
|
110
|
+
pass
|
|
111
|
+
return asyncio.run(fn(*args, **kwargs))
|
|
112
|
+
return sync_wrapper
|
|
113
|
+
if inspect.iscoroutinefunction(fn):
|
|
114
|
+
return fn
|
|
115
|
+
async def async_wrapper(*args, **kwargs):
|
|
116
|
+
return fn(*args, **kwargs)
|
|
117
|
+
return async_wrapper
|
|
118
|
+
|
|
119
|
+
def build_edges(node_names: Sequence[str], edges: Sequence[Any]):
|
|
120
|
+
"""Return (regular_edges, conditional_edges) normalized + auto START/END guarded."""
|
|
121
|
+
all_nodes = set(node_names)
|
|
122
|
+
regular_edges, conditional_edges = [], []
|
|
123
|
+
for edge in edges:
|
|
124
|
+
if isinstance(edge, Edge):
|
|
125
|
+
regular_edges.append((edge.start, edge.end))
|
|
126
|
+
elif isinstance(edge, ConditionalEdge):
|
|
127
|
+
for target in edge.targets:
|
|
128
|
+
if target not in all_nodes and target not in (START, END):
|
|
129
|
+
raise ValueError(f"ConditionalEdge target '{target}' is not a known node or START/END")
|
|
130
|
+
conditional_edges.append((edge.start, make_router_wrapper(edge.router_fn, edge.targets), {t: t for t in edge.targets}))
|
|
131
|
+
else:
|
|
132
|
+
raise ValueError(f"Unknown edge type: {edge}")
|
|
133
|
+
if regular_edges and not any(s == START for s, _ in regular_edges):
|
|
134
|
+
regular_edges = [(START, regular_edges[0][0]), *regular_edges]
|
|
135
|
+
if regular_edges and not any(e == END for _, e in regular_edges):
|
|
136
|
+
regular_edges = [*regular_edges, (regular_edges[-1][1], END)]
|
|
137
|
+
validate_edges(regular_edges, all_nodes)
|
|
138
|
+
validate_conditional_edges(conditional_edges, all_nodes)
|
|
139
|
+
return regular_edges, conditional_edges
|
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
from ai_infra.llm.core import CoreLLM, CoreAgent, BaseLLMCore
|
|
2
|
+
from ai_infra.llm.utils.settings import ModelSettings
|
|
3
|
+
from ai_infra.llm.providers import Providers
|
|
4
|
+
from ai_infra.llm.providers.models import Models
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
__all__ = [
|
|
8
|
+
"CoreLLM",
|
|
9
|
+
"CoreAgent",
|
|
10
|
+
"ModelSettings",
|
|
11
|
+
"Models",
|
|
12
|
+
"Providers",
|
|
13
|
+
]
|