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.
Files changed (80) hide show
  1. ai_infra-0.1.3/PKG-INFO +39 -0
  2. ai_infra-0.1.3/README.md +1 -0
  3. ai_infra-0.1.3/pyproject.toml +59 -0
  4. ai_infra-0.1.3/src/ai_infra/__init__.py +20 -0
  5. ai_infra-0.1.3/src/ai_infra/graph/__init__.py +8 -0
  6. ai_infra-0.1.3/src/ai_infra/graph/core.py +156 -0
  7. ai_infra-0.1.3/src/ai_infra/graph/examples/01_graph_basic.py +45 -0
  8. ai_infra-0.1.3/src/ai_infra/graph/examples/02_graph_stream_values.py +27 -0
  9. ai_infra-0.1.3/src/ai_infra/graph/examples/__init__.py +0 -0
  10. ai_infra-0.1.3/src/ai_infra/graph/models.py +35 -0
  11. ai_infra-0.1.3/src/ai_infra/graph/utils.py +139 -0
  12. ai_infra-0.1.3/src/ai_infra/llm/__init__.py +13 -0
  13. ai_infra-0.1.3/src/ai_infra/llm/core.py +357 -0
  14. ai_infra-0.1.3/src/ai_infra/llm/examples/01_agent_basic.py +16 -0
  15. ai_infra-0.1.3/src/ai_infra/llm/examples/02_llm_chat_basic.py +16 -0
  16. ai_infra-0.1.3/src/ai_infra/llm/examples/03_structured_output.py +27 -0
  17. ai_infra-0.1.3/src/ai_infra/llm/examples/04_agent_stream.py +29 -0
  18. ai_infra-0.1.3/src/ai_infra/llm/examples/05_tool_controls.py +34 -0
  19. ai_infra-0.1.3/src/ai_infra/llm/examples/06_hitl.py +43 -0
  20. ai_infra-0.1.3/src/ai_infra/llm/examples/07_retry.py +27 -0
  21. ai_infra-0.1.3/src/ai_infra/llm/examples/08_agent_stream_tokens.py +22 -0
  22. ai_infra-0.1.3/src/ai_infra/llm/examples/09_chat_stream.py +25 -0
  23. ai_infra-0.1.3/src/ai_infra/llm/examples/__init__.py +0 -0
  24. ai_infra-0.1.3/src/ai_infra/llm/providers/__init__.py +2 -0
  25. ai_infra-0.1.3/src/ai_infra/llm/providers/models.py +28 -0
  26. ai_infra-0.1.3/src/ai_infra/llm/providers/providers.py +5 -0
  27. ai_infra-0.1.3/src/ai_infra/llm/tools/__init__.py +2 -0
  28. ai_infra-0.1.3/src/ai_infra/llm/tools/tool_controls.py +114 -0
  29. ai_infra-0.1.3/src/ai_infra/llm/tools/tools.py +215 -0
  30. ai_infra-0.1.3/src/ai_infra/llm/utils/__init__.py +38 -0
  31. ai_infra-0.1.3/src/ai_infra/llm/utils/fallbacks.py +115 -0
  32. ai_infra-0.1.3/src/ai_infra/llm/utils/messages.py +24 -0
  33. ai_infra-0.1.3/src/ai_infra/llm/utils/model_init.py +24 -0
  34. ai_infra-0.1.3/src/ai_infra/llm/utils/retry.py +16 -0
  35. ai_infra-0.1.3/src/ai_infra/llm/utils/runtime_bind.py +169 -0
  36. ai_infra-0.1.3/src/ai_infra/llm/utils/settings.py +9 -0
  37. ai_infra-0.1.3/src/ai_infra/llm/utils/validation.py +23 -0
  38. ai_infra-0.1.3/src/ai_infra/mcp/__init__.py +7 -0
  39. ai_infra-0.1.3/src/ai_infra/mcp/client/__init__.py +0 -0
  40. ai_infra-0.1.3/src/ai_infra/mcp/client/core.py +460 -0
  41. ai_infra-0.1.3/src/ai_infra/mcp/client/models.py +29 -0
  42. ai_infra-0.1.3/src/ai_infra/mcp/examples/01_mcps.py +9 -0
  43. ai_infra-0.1.3/src/ai_infra/mcp/examples/__init__.py +0 -0
  44. ai_infra-0.1.3/src/ai_infra/mcp/examples/agents/01_streamable_http_agent.py +28 -0
  45. ai_infra-0.1.3/src/ai_infra/mcp/examples/agents/02_multi_server_agent.py +32 -0
  46. ai_infra-0.1.3/src/ai_infra/mcp/examples/agents/__init__.py +0 -0
  47. ai_infra-0.1.3/src/ai_infra/mcp/examples/client/01_sse.py +18 -0
  48. ai_infra-0.1.3/src/ai_infra/mcp/examples/client/02_stdio.py +20 -0
  49. ai_infra-0.1.3/src/ai_infra/mcp/examples/client/03_streamable_http.py +22 -0
  50. ai_infra-0.1.3/src/ai_infra/mcp/examples/client/04_stdio.py +23 -0
  51. ai_infra-0.1.3/src/ai_infra/mcp/examples/client/05_openapi.py +15 -0
  52. ai_infra-0.1.3/src/ai_infra/mcp/examples/client/06_multi_server_client.py +20 -0
  53. ai_infra-0.1.3/src/ai_infra/mcp/examples/client/06_server_metadata_from_client.py +30 -0
  54. ai_infra-0.1.3/src/ai_infra/mcp/examples/client/__init__.py +0 -0
  55. ai_infra-0.1.3/src/ai_infra/mcp/examples/resources/apiframeworks.json +3 -0
  56. ai_infra-0.1.3/src/ai_infra/mcp/examples/resources/spotify.yaml +6954 -0
  57. ai_infra-0.1.3/src/ai_infra/mcp/examples/server/01_sse.py +10 -0
  58. ai_infra-0.1.3/src/ai_infra/mcp/examples/server/02_stdio.py +12 -0
  59. ai_infra-0.1.3/src/ai_infra/mcp/examples/server/03_streamable_http.py +11 -0
  60. ai_infra-0.1.3/src/ai_infra/mcp/examples/server/04_openapi.py +29 -0
  61. ai_infra-0.1.3/src/ai_infra/mcp/examples/server/__init__.py +0 -0
  62. ai_infra-0.1.3/src/ai_infra/mcp/examples/server/fastapi/01_add_app.py +38 -0
  63. ai_infra-0.1.3/src/ai_infra/mcp/examples/server/fastapi/02_add_fastmcp.py +32 -0
  64. ai_infra-0.1.3/src/ai_infra/mcp/examples/server/fastapi/03_raw_mount.py +29 -0
  65. ai_infra-0.1.3/src/ai_infra/mcp/examples/server/fastapi/__init__.py +0 -0
  66. ai_infra-0.1.3/src/ai_infra/mcp/examples/server/fastapi/from_module/01_from_module_fastmcp.py +9 -0
  67. ai_infra-0.1.3/src/ai_infra/mcp/examples/server/fastapi/from_module/02_from_modle_asgi.py +13 -0
  68. ai_infra-0.1.3/src/ai_infra/mcp/examples/server/fastapi/from_module/03_from_module.py +22 -0
  69. ai_infra-0.1.3/src/ai_infra/mcp/examples/server/fastapi/from_module/__init__.py +0 -0
  70. ai_infra-0.1.3/src/ai_infra/mcp/server/__init__.py +0 -0
  71. ai_infra-0.1.3/src/ai_infra/mcp/server/core.py +379 -0
  72. ai_infra-0.1.3/src/ai_infra/mcp/server/models.py +12 -0
  73. ai_infra-0.1.3/src/ai_infra/mcp/server/openapi/__init__.py +4 -0
  74. ai_infra-0.1.3/src/ai_infra/mcp/server/openapi/builder.py +582 -0
  75. ai_infra-0.1.3/src/ai_infra/mcp/server/openapi/constants.py +3 -0
  76. ai_infra-0.1.3/src/ai_infra/mcp/server/openapi/io.py +19 -0
  77. ai_infra-0.1.3/src/ai_infra/mcp/server/openapi/models.py +60 -0
  78. ai_infra-0.1.3/src/ai_infra/mcp/server/openapi/runtime.py +90 -0
  79. ai_infra-0.1.3/src/ai_infra/mcp/server/tools.py +56 -0
  80. ai_infra-0.1.3/src/ai_infra/py.typed +0 -0
@@ -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
+
@@ -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,8 @@
1
+ from ai_infra.graph.models import Edge, ConditionalEdge
2
+ from ai_infra.graph.core import CoreGraph
3
+
4
+ __all__ = [
5
+ "CoreGraph",
6
+ "Edge",
7
+ "ConditionalEdge",
8
+ ]
@@ -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
+ ]