Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
200 changes: 150 additions & 50 deletions src/fenn/agents/__init__.py
Original file line number Diff line number Diff line change
@@ -1,124 +1,224 @@
import asyncio, warnings, copy, time
import asyncio
import copy
import time
import warnings

_TERMINAL = object() # sentinel for explicit terminal transitions in Flow.connect


class BaseNode:
def __init__(self): self.params, self.successors = {}, {}
def set_params(self, params): self.params = params
def prep(self, shared): pass
def exec(self, prep_res): pass
def post(self, shared, prep_res, exec_res): pass
def _exec(self, prep_res): return self.exec(prep_res)
def _run(self, shared): p = self.prep(shared); e = self._exec(p); return self.post(shared, p, e)
def __init__(self):
self.params, self.successors = {}, {}

def set_params(self, params):
self.params = params

def prep(self, shared):
pass

def exec(self, prep_res):
pass

def post(self, shared, prep_res, exec_res):
pass

def _exec(self, prep_res):
return self.exec(prep_res)

def _run(self, shared):
p = self.prep(shared)
e = self._exec(p)
return self.post(shared, p, e)

def run(self, shared):
if self.successors: warnings.warn("Node won't run successors. Use Flow.")
if self.successors:
warnings.warn("Node won't run successors. Use Flow.")
return self._run(shared)


class Node(BaseNode):
def __init__(self, max_retries=1, wait=0): super().__init__(); self.max_retries, self.wait = max_retries, wait
def exec_fallback(self, prep_res, exc): raise exc
def __init__(self, max_retries=1, wait=0):
super().__init__()
self.max_retries, self.wait = max_retries, wait

def exec_fallback(self, prep_res, exc):
raise exc

def _exec(self, prep_res):
for self.cur_retry in range(self.max_retries):
try: return self.exec(prep_res)
try:
return self.exec(prep_res)
except Exception as e:
if self.cur_retry == self.max_retries - 1: return self.exec_fallback(prep_res, e)
if self.wait > 0: time.sleep(self.wait)
if self.cur_retry == self.max_retries - 1:
return self.exec_fallback(prep_res, e)
if self.wait > 0:
time.sleep(self.wait)


class BatchNode(Node):
def _exec(self, items): return [super(BatchNode, self)._exec(i) for i in (items or [])]
def _exec(self, items):
return [super(BatchNode, self)._exec(i) for i in (items or [])]


class Flow(BaseNode):
def __init__(self, start=None): super().__init__(); self.start_node = start
def __init__(self, start=None):
super().__init__()
self.start_node = start

def start(self, start): self.start_node = start; return start
def start(self, start):
self.start_node = start
return start

def connect(self, src, dst, action="default"):
if action in src.successors: warnings.warn(f"Overwriting successor for action '{action}'")
if action in src.successors:
warnings.warn(f"Overwriting successor for action '{action}'")
src.successors[action] = _TERMINAL if dst is None else dst
return self

def get_next_node(self, curr, action):
nxt = curr.successors.get(action or "default")
if nxt is _TERMINAL: return None
if nxt is None and curr.successors: warnings.warn(f"Flow ends: '{action}' not found in {list(curr.successors)}")
if nxt is _TERMINAL:
return None
if nxt is None and curr.successors:
warnings.warn(f"Flow ends: '{action}' not found in {list(curr.successors)}")
return nxt

def _orch(self, shared, params=None):
curr, p, last_action = copy.copy(self.start_node), (params or {**self.params}), None
while curr: curr.set_params(p); last_action = curr._run(shared); curr = copy.copy(self.get_next_node(curr, last_action))
curr, p, last_action = (
copy.copy(self.start_node),
(params or {**self.params}),
None,
)
while curr:
curr.set_params(p)
last_action = curr._run(shared)
curr = copy.copy(self.get_next_node(curr, last_action))
return last_action

def _run(self, shared): p = self.prep(shared); o = self._orch(shared); return self.post(shared, p, o)
def post(self, shared, prep_res, exec_res): return exec_res
def _run(self, shared):
p = self.prep(shared)
o = self._orch(shared)
return self.post(shared, p, o)

def post(self, shared, prep_res, exec_res):
return exec_res


class BatchFlow(Flow):
def _run(self, shared):
pr = self.prep(shared) or []
for bp in pr: self._orch(shared, {**self.params, **bp})
for bp in pr:
self._orch(shared, {**self.params, **bp})
return self.post(shared, pr, None)


class AsyncNode(Node):
async def prep_async(self, shared): pass
async def exec_async(self, prep_res): pass
async def exec_fallback_async(self, prep_res, exc): raise exc
async def post_async(self, shared, prep_res, exec_res): pass
async def prep_async(self, shared):
pass

async def exec_async(self, prep_res):
pass

async def exec_fallback_async(self, prep_res, exc):
raise exc

async def post_async(self, shared, prep_res, exec_res):
pass

async def _exec(self, prep_res):
for self.cur_retry in range(self.max_retries):
try: return await self.exec_async(prep_res)
try:
return await self.exec_async(prep_res)
except Exception as e:
if self.cur_retry == self.max_retries - 1: return await self.exec_fallback_async(prep_res, e)
if self.wait > 0: await asyncio.sleep(self.wait)
if self.cur_retry == self.max_retries - 1:
return await self.exec_fallback_async(prep_res, e)
if self.wait > 0:
await asyncio.sleep(self.wait)

async def run_async(self, shared):
if self.successors: warnings.warn("Node won't run successors. Use AsyncFlow.")
if self.successors:
warnings.warn("Node won't run successors. Use AsyncFlow.")
return await self._run_async(shared)
async def _run_async(self, shared): p = await self.prep_async(shared); e = await self._exec(p); return await self.post_async(shared, p, e)
def _run(self, shared): raise RuntimeError("Use run_async.")

async def _run_async(self, shared):
p = await self.prep_async(shared)
e = await self._exec(p)
return await self.post_async(shared, p, e)

def _run(self, shared):
raise RuntimeError("Use run_async.")


class AsyncBatchNode(AsyncNode, BatchNode):
async def _exec(self, items): return [await super(AsyncBatchNode, self)._exec(i) for i in items]
async def _exec(self, items):
return [await super(AsyncBatchNode, self)._exec(i) for i in items]


class AsyncParallelBatchNode(AsyncNode, BatchNode):
async def _exec(self, items): return await asyncio.gather(*(super(AsyncParallelBatchNode, self)._exec(i) for i in items))
async def _exec(self, items):
return await asyncio.gather(
*(super(AsyncParallelBatchNode, self)._exec(i) for i in items)
)


class AsyncFlow(Flow, AsyncNode):
async def _orch_async(self, shared, params=None):
curr, p, last_action = copy.copy(self.start_node), (params or {**self.params}), None
while curr: curr.set_params(p); last_action = await curr._run_async(shared) if isinstance(curr, AsyncNode) else curr._run(shared); curr = copy.copy(self.get_next_node(curr, last_action))
curr, p, last_action = (
copy.copy(self.start_node),
(params or {**self.params}),
None,
)
while curr:
curr.set_params(p)
last_action = (
await curr._run_async(shared)
if isinstance(curr, AsyncNode)
else curr._run(shared)
)
curr = copy.copy(self.get_next_node(curr, last_action))
return last_action
async def _run_async(self, shared): p = await self.prep_async(shared); o = await self._orch_async(shared); return await self.post_async(shared, p, o)
async def post_async(self, shared, prep_res, exec_res): return exec_res

async def _run_async(self, shared):
p = await self.prep_async(shared)
o = await self._orch_async(shared)
return await self.post_async(shared, p, o)

async def post_async(self, shared, prep_res, exec_res):
return exec_res


class AsyncBatchFlow(AsyncFlow, BatchFlow):
async def _run_async(self, shared):
pr = await self.prep_async(shared) or []
for bp in pr: await self._orch_async(shared, {**self.params, **bp})
for bp in pr:
await self._orch_async(shared, {**self.params, **bp})
return await self.post_async(shared, pr, None)


class AsyncParallelBatchFlow(AsyncFlow, BatchFlow):
async def _run_async(self, shared):
pr = await self.prep_async(shared) or []
await asyncio.gather(*(self._orch_async(shared, {**self.params, **bp}) for bp in pr))
await asyncio.gather(
*(self._orch_async(shared, {**self.params, **bp}) for bp in pr)
)
return await self.post_async(shared, pr, None)


from .llm import LLMClient
from .rag import RAGNode
from .llm import LLMClient # noqa: E402 - avoid circular import with .llm
from .rag import RAGNode # noqa: E402 - avoid circular import with .rag

__all__ = [
"BaseNode", "Node", "BatchNode",
"Flow", "BatchFlow",
"AsyncNode", "AsyncBatchNode", "AsyncParallelBatchNode",
"AsyncFlow", "AsyncBatchFlow", "AsyncParallelBatchFlow",
"LLMClient", "RAGNode",
"BaseNode",
"Node",
"BatchNode",
"Flow",
"BatchFlow",
"AsyncNode",
"AsyncBatchNode",
"AsyncParallelBatchNode",
"AsyncFlow",
"AsyncBatchFlow",
"AsyncParallelBatchFlow",
"LLMClient",
"RAGNode",
]
59 changes: 42 additions & 17 deletions src/fenn/agents/__init__.pyi
Original file line number Diff line number Diff line change
@@ -1,9 +1,8 @@
import asyncio
from typing import Any, Dict, Iterator, List, Optional, Union, TypeVar, Generic
from typing import Any, Dict, Generic, Iterator, List, Optional, TypeVar, Union

_PrepResult = TypeVar('_PrepResult')
_ExecResult = TypeVar('_ExecResult')
_PostResult = TypeVar('_PostResult')
_PrepResult = TypeVar("_PrepResult")
_ExecResult = TypeVar("_ExecResult")
_PostResult = TypeVar("_PostResult")

ParamValue = Union[str, int, float, bool, None, List[Any], Dict[str, Any]]
SharedData = Dict[str, Any]
Expand All @@ -18,7 +17,9 @@ class BaseNode(Generic[_PrepResult, _ExecResult, _PostResult]):
def set_params(self, params: Params) -> None: ...
def prep(self, shared: SharedData) -> _PrepResult: ...
def exec(self, prep_res: _PrepResult) -> _ExecResult: ...
def post(self, shared: SharedData, prep_res: _PrepResult, exec_res: _ExecResult) -> _PostResult: ...
def post(
self, shared: SharedData, prep_res: _PrepResult, exec_res: _ExecResult
) -> _PostResult: ...
def _exec(self, prep_res: _PrepResult) -> _ExecResult: ...
def _run(self, shared: SharedData) -> _PostResult: ...
def run(self, shared: SharedData) -> _PostResult: ...
Expand Down Expand Up @@ -51,36 +52,60 @@ class Flow(BaseNode[_PrepResult, Any, _PostResult]):
) -> Optional[BaseNode[Any, Any, Any]]: ...
def _orch(self, shared: SharedData, params: Optional[Params] = None) -> Any: ...
def _run(self, shared: SharedData) -> _PostResult: ...
def post(self, shared: SharedData, prep_res: _PrepResult, exec_res: Any) -> _PostResult: ...
def post(
self, shared: SharedData, prep_res: _PrepResult, exec_res: Any
) -> _PostResult: ...

class BatchFlow(Flow[Optional[List[Params]], Any, _PostResult]):
def _run(self, shared: SharedData) -> _PostResult: ...

class AsyncNode(Node[_PrepResult, _ExecResult, _PostResult]):
async def prep_async(self, shared: SharedData) -> _PrepResult: ...
async def exec_async(self, prep_res: _PrepResult) -> _ExecResult: ...
async def exec_fallback_async(self, prep_res: _PrepResult, exc: Exception) -> _ExecResult: ...
async def post_async(self, shared: SharedData, prep_res: _PrepResult, exec_res: _ExecResult) -> _PostResult: ...
async def exec_fallback_async(
self, prep_res: _PrepResult, exc: Exception
) -> _ExecResult: ...
async def post_async(
self, shared: SharedData, prep_res: _PrepResult, exec_res: _ExecResult
) -> _PostResult: ...
async def _exec(self, prep_res: _PrepResult) -> _ExecResult: ...
async def run_async(self, shared: SharedData) -> _PostResult: ...
async def _run_async(self, shared: SharedData) -> _PostResult: ...
def _run(self, shared: SharedData) -> _PostResult: ...

class AsyncBatchNode(AsyncNode[Optional[List[_PrepResult]], List[_ExecResult], _PostResult], BatchNode[Optional[List[_PrepResult]], List[_ExecResult], _PostResult]):
class AsyncBatchNode(
AsyncNode[Optional[List[_PrepResult]], List[_ExecResult], _PostResult],
BatchNode[Optional[List[_PrepResult]], List[_ExecResult], _PostResult],
):
async def _exec(self, items: Optional[List[_PrepResult]]) -> List[_ExecResult]: ...

class AsyncParallelBatchNode(AsyncNode[Optional[List[_PrepResult]], List[_ExecResult], _PostResult], BatchNode[Optional[List[_PrepResult]], List[_ExecResult], _PostResult]):
class AsyncParallelBatchNode(
AsyncNode[Optional[List[_PrepResult]], List[_ExecResult], _PostResult],
BatchNode[Optional[List[_PrepResult]], List[_ExecResult], _PostResult],
):
async def _exec(self, items: Optional[List[_PrepResult]]) -> List[_ExecResult]: ...

class AsyncFlow(Flow[_PrepResult, Any, _PostResult], AsyncNode[_PrepResult, Any, _PostResult]):
async def _orch_async(self, shared: SharedData, params: Optional[Params] = None) -> Any: ...
class AsyncFlow(
Flow[_PrepResult, Any, _PostResult], AsyncNode[_PrepResult, Any, _PostResult]
):
async def _orch_async(
self, shared: SharedData, params: Optional[Params] = None
) -> Any: ...
async def _run_async(self, shared: SharedData) -> _PostResult: ...
async def post_async(self, shared: SharedData, prep_res: _PrepResult, exec_res: Any) -> _PostResult: ...

class AsyncBatchFlow(AsyncFlow[Optional[List[Params]], Any, _PostResult], BatchFlow[Optional[List[Params]], Any, _PostResult]):
async def post_async(
self, shared: SharedData, prep_res: _PrepResult, exec_res: Any
) -> _PostResult: ...

class AsyncBatchFlow(
AsyncFlow[Optional[List[Params]], Any, _PostResult],
BatchFlow[Optional[List[Params]], Any, _PostResult],
):
async def _run_async(self, shared: SharedData) -> _PostResult: ...

class AsyncParallelBatchFlow(AsyncFlow[Optional[List[Params]], Any, _PostResult], BatchFlow[Optional[List[Params]], Any, _PostResult]):
class AsyncParallelBatchFlow(
AsyncFlow[Optional[List[Params]], Any, _PostResult],
BatchFlow[Optional[List[Params]], Any, _PostResult],
):
async def _run_async(self, shared: SharedData) -> _PostResult: ...

class LLMClient:
Expand Down
Loading
Loading