mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-11 03:40:03 +00:00
refactor(core): update BaseOp to improve sub-ops handling and tool call management
This commit is contained in:
parent
472e069bc5
commit
566a773591
6 changed files with 395 additions and 48 deletions
|
|
@ -27,11 +27,8 @@ class RuntimeContext(BaseContext):
|
|||
if context is None:
|
||||
return cls(**kwargs)
|
||||
|
||||
new_instance = cls(response=context.response, stream_queue=context.stream_queue)
|
||||
new_instance.update(context)
|
||||
if kwargs:
|
||||
new_instance.update(kwargs)
|
||||
return new_instance
|
||||
context.update(kwargs)
|
||||
return context
|
||||
|
||||
async def _enqueue(self, chunk: StreamChunk) -> None:
|
||||
"""Internal helper to put a chunk into the queue if it exists."""
|
||||
|
|
|
|||
|
|
@ -4,15 +4,15 @@ import asyncio
|
|||
import copy
|
||||
import inspect
|
||||
from pathlib import Path
|
||||
from typing import Callable, Any, Union
|
||||
from typing import Callable, Any, Optional
|
||||
|
||||
from loguru import logger
|
||||
from tqdm import tqdm
|
||||
|
||||
from ..context import RuntimeContext, PromptHandler, C, BaseContext
|
||||
from ..context import RuntimeContext, PromptHandler, C
|
||||
from ..embedding import BaseEmbeddingModel
|
||||
from ..llm import BaseLLM
|
||||
from ..schema import ToolCall, ToolAttr
|
||||
from ..schema import ToolCall, ToolAttr, Response
|
||||
from ..token_counter import BaseTokenCounter
|
||||
from ..utils import camel_to_snake, CacheHandler, timer
|
||||
from ..vector_store import BaseVectorStore
|
||||
|
|
@ -40,10 +40,10 @@ class BaseOp:
|
|||
token_counter: str | BaseTokenCounter = "default",
|
||||
enable_cache: bool = False,
|
||||
cache_path: str = "cache/op",
|
||||
sub_ops: Union[list["BaseOp"], dict[str, "BaseOp"], "BaseOp", None] = None,
|
||||
sub_ops: dict[str, "BaseOp"] | list["BaseOp"] | Optional["BaseOp"] = None,
|
||||
input_mapping: dict[str, str] | None = None,
|
||||
output_mapping: dict[str, str] | None = None,
|
||||
enable_tool_response: bool = False,
|
||||
save_response_result: bool = False,
|
||||
enable_sync_thread_pool: bool = True,
|
||||
max_retries: int = 1,
|
||||
raise_exception: bool = False,
|
||||
|
|
@ -62,12 +62,12 @@ class BaseOp:
|
|||
|
||||
self.enable_cache = enable_cache
|
||||
self.cache_path = cache_path
|
||||
self.sub_ops = BaseContext[str, BaseOp]()
|
||||
self.sub_ops: list[BaseOp] = []
|
||||
self.add_sub_ops(sub_ops)
|
||||
|
||||
self.input_mapping = input_mapping
|
||||
self.output_mapping = output_mapping
|
||||
self.enable_tool_response = enable_tool_response
|
||||
self.save_response_result = save_response_result
|
||||
self.enable_sync_thread_pool = enable_sync_thread_pool
|
||||
self.max_retries = max(1, max_retries)
|
||||
self.raise_exception = raise_exception
|
||||
|
|
@ -89,7 +89,7 @@ class BaseOp:
|
|||
|
||||
def _validate_inputs(self):
|
||||
"""Ensure all required tool inputs are present in context."""
|
||||
if self.tool_call:
|
||||
if self.tool_call is not None:
|
||||
parameters = self.tool_call.parameters
|
||||
if parameters.type == "object" and parameters.properties:
|
||||
required_list = parameters.required or []
|
||||
|
|
@ -98,30 +98,47 @@ class BaseOp:
|
|||
|
||||
def _handle_failure(self, e: Exception, attempt: int):
|
||||
"""Log failures and handle final retry logic."""
|
||||
logger.exception(f"{self.name} failed (attempt {attempt + 1}): {e}")
|
||||
message = f"{self.name} failed (attempt {attempt + 1}): {e}"
|
||||
if attempt == self.max_retries - 1:
|
||||
logger.exception(message)
|
||||
if self.raise_exception:
|
||||
raise e
|
||||
|
||||
if self.tool_call:
|
||||
if self.tool_call is not None:
|
||||
self.output = f"{self.name} failed: {e}"
|
||||
else:
|
||||
logger.warning(message)
|
||||
|
||||
@property
|
||||
def tool_call(self) -> ToolCall:
|
||||
def tool_call(self) -> ToolCall | None:
|
||||
"""Lazily construct and return the tool call metadata."""
|
||||
if self._tool_call is None:
|
||||
self._tool_call = self._build_tool_call()
|
||||
assert self._tool_call, "tool_call is not defined!"
|
||||
if self._tool_call is None:
|
||||
return None
|
||||
|
||||
self._tool_call.name = self._tool_call.name or self.name
|
||||
if not self._tool_call.output.properties:
|
||||
self._tool_call.output = ToolAttr(
|
||||
type="object",
|
||||
properties={
|
||||
f"{self.name}_result": ToolAttr(type="string", description=f"Execution result of {self.name}"),
|
||||
},
|
||||
)
|
||||
self._tool_call.output.properties = {
|
||||
f"{self.name}_result": ToolAttr(type="string", description=f"Execution result of {self.name}"),
|
||||
}
|
||||
return self._tool_call
|
||||
|
||||
def set_tool_call(self, tool_call: ToolCall | dict):
|
||||
"""Set the tool call."""
|
||||
if isinstance(tool_call, dict):
|
||||
self._tool_call = ToolCall(**tool_call)
|
||||
elif isinstance(tool_call, ToolCall):
|
||||
self._tool_call = tool_call
|
||||
else:
|
||||
raise ValueError(f"Invalid tool call: {tool_call}")
|
||||
|
||||
self._tool_call.name = self._tool_call.name or self.name
|
||||
if not self._tool_call.output.properties:
|
||||
self._tool_call.output.properties = {
|
||||
f"{self.name}_result": ToolAttr(type="string", description=f"Execution result of {self.name}"),
|
||||
}
|
||||
|
||||
@property
|
||||
def input_dict(self) -> dict:
|
||||
"""Extract required and optional inputs from context based on schema."""
|
||||
|
|
@ -137,6 +154,7 @@ class BaseOp:
|
|||
output_properties = self.tool_call.output.properties
|
||||
if not output_properties:
|
||||
return None
|
||||
|
||||
keys = list(output_properties.keys())
|
||||
return self.context[keys[0]]
|
||||
|
||||
|
|
@ -146,6 +164,7 @@ class BaseOp:
|
|||
output_properties = self.tool_call.output.properties
|
||||
if not output_properties:
|
||||
return
|
||||
|
||||
keys = list(output_properties.keys())
|
||||
self.context[keys[0]] = value
|
||||
|
||||
|
|
@ -188,10 +207,7 @@ class BaseOp:
|
|||
"""Lazily initialize and return the token counter instance."""
|
||||
if isinstance(self._token_counter, str):
|
||||
cfg = C.service_config.token_counter[self._token_counter]
|
||||
self._token_counter = C.get_token_counter_class(cfg.backend)(
|
||||
model_name=cfg.model_name,
|
||||
**cfg.model_extra,
|
||||
)
|
||||
self._token_counter = C.get_token_counter_class(cfg.backend)(model_name=cfg.model_name, **cfg.model_extra)
|
||||
return self._token_counter
|
||||
|
||||
@property
|
||||
|
|
@ -199,6 +215,11 @@ class BaseOp:
|
|||
"""Get service configuration metadata."""
|
||||
return C.service_config.model_extra
|
||||
|
||||
@property
|
||||
def response(self) -> Response:
|
||||
"""Get the response object."""
|
||||
return self.context.response
|
||||
|
||||
async def before_execute(self):
|
||||
"""Prepare context and validate before async execution."""
|
||||
self.context.apply_mapping(self.input_mapping)
|
||||
|
|
@ -210,7 +231,7 @@ class BaseOp:
|
|||
async def after_execute(self):
|
||||
"""Finalize context and mappings after async execution."""
|
||||
self.context.apply_mapping(self.output_mapping)
|
||||
if self.tool_call and self.enable_tool_response:
|
||||
if self.tool_call is not None and self.save_response_result:
|
||||
self.context.response.answer = self.output
|
||||
|
||||
if not isinstance(self._llm, str) and hasattr(self._llm, "close"):
|
||||
|
|
@ -229,7 +250,7 @@ class BaseOp:
|
|||
def after_execute_sync(self):
|
||||
"""Finalize context and mappings after sync execution."""
|
||||
self.context.apply_mapping(self.output_mapping)
|
||||
if self.tool_call and self.enable_tool_response:
|
||||
if self.tool_call is not None and self.save_response_result:
|
||||
self.context.response.answer = self.output
|
||||
|
||||
if not isinstance(self._llm, str) and hasattr(self._llm, "close_sync"):
|
||||
|
|
@ -249,7 +270,7 @@ class BaseOp:
|
|||
break
|
||||
except Exception as e:
|
||||
self._handle_failure(e, i)
|
||||
return self.output if self.tool_call else None
|
||||
return self.output if self.tool_call is not None else None
|
||||
|
||||
async def call(self, context: RuntimeContext = None, **kwargs):
|
||||
"""Execute the operator asynchronously with retry logic."""
|
||||
|
|
@ -262,7 +283,7 @@ class BaseOp:
|
|||
break
|
||||
except Exception as e:
|
||||
self._handle_failure(e, i)
|
||||
return self.output if self.tool_call else None
|
||||
return self.output if self.tool_call is not None else None
|
||||
|
||||
def submit_sync_task(self, fn: Callable, *args, **kwargs) -> "BaseOp":
|
||||
"""Submit a task to the thread pool or local queue."""
|
||||
|
|
@ -301,23 +322,27 @@ class BaseOp:
|
|||
finally:
|
||||
self._pending_tasks.clear()
|
||||
|
||||
def add_sub_ops(self, sub_ops: Union[list["BaseOp"], dict[str, "BaseOp"], "BaseOp", None]):
|
||||
"""Add child operators to this operator's sub_ops context."""
|
||||
def add_sub_ops(self, sub_ops: dict[str, "BaseOp"] | list["BaseOp"] | Optional["BaseOp"]):
|
||||
"""Add child operators to this operator's sub_ops."""
|
||||
if not sub_ops:
|
||||
return
|
||||
|
||||
if isinstance(sub_ops, dict):
|
||||
ops_dict = sub_ops
|
||||
for name, op in sub_ops.items():
|
||||
assert self.async_mode == op.async_mode, "Async mode mismatch!"
|
||||
op.name = name
|
||||
self.sub_ops.append(op)
|
||||
elif isinstance(sub_ops, list):
|
||||
for op in sub_ops:
|
||||
assert self.async_mode == op.async_mode, "Async mode mismatch!"
|
||||
self.sub_ops.append(op)
|
||||
else:
|
||||
ops_dict = {op.name: op for op in (sub_ops if isinstance(sub_ops, list) else [sub_ops])}
|
||||
|
||||
for name, op in ops_dict.items():
|
||||
assert self.async_mode == op.async_mode, "Async mode mismatch!"
|
||||
self.sub_ops[name] = op
|
||||
assert self.async_mode == sub_ops.async_mode, "Async mode mismatch!"
|
||||
self.sub_ops.append(sub_ops)
|
||||
|
||||
def add_sub_op(self, sub_op: "BaseOp"):
|
||||
"""Add a single child operator to this operator's sub_ops context."""
|
||||
self.add_sub_ops(sub_op)
|
||||
"""Add a single child operator to this operator's sub_ops."""
|
||||
self.sub_ops.append(sub_op)
|
||||
|
||||
def __lshift__(self, ops):
|
||||
"""Operator overload for adding sub-operators."""
|
||||
|
|
@ -350,7 +375,7 @@ class BaseOp:
|
|||
|
||||
def copy(self, **kwargs):
|
||||
"""Create a copy of this operator with optional parameter overrides."""
|
||||
copy_op = self.__class__(*self._init_args, **self._init_kwargs, **kwargs)
|
||||
copy_op = self.__class__(*self._init_args, **{**self._init_kwargs, **kwargs})
|
||||
if self.sub_ops:
|
||||
copy_op.sub_ops.clear()
|
||||
copy_op.add_sub_ops(self.sub_ops)
|
||||
|
|
|
|||
|
|
@ -8,14 +8,14 @@ class ParallelOp(BaseOp):
|
|||
|
||||
async def execute(self):
|
||||
"""Executes all sub-operations concurrently using asynchronous tasks."""
|
||||
for op in self.sub_ops.values():
|
||||
for op in self.sub_ops:
|
||||
assert op.async_mode
|
||||
self.submit_async_task(op.call, context=self.context)
|
||||
await self.join_async_tasks()
|
||||
|
||||
def execute_sync(self):
|
||||
"""Executes all sub-operations concurrently using synchronous task management."""
|
||||
for op in self.sub_ops.values():
|
||||
for op in self.sub_ops:
|
||||
assert not op.async_mode
|
||||
self.submit_sync_task(op.call_sync, context=self.context)
|
||||
self.join_sync_tasks()
|
||||
|
|
|
|||
|
|
@ -8,13 +8,13 @@ class SequentialOp(BaseOp):
|
|||
|
||||
async def execute(self):
|
||||
"""Executes sub-operations sequentially using asynchronous awaits."""
|
||||
for op in self.sub_ops.values():
|
||||
for op in self.sub_ops:
|
||||
assert op.async_mode
|
||||
await op.call(context=self.context)
|
||||
|
||||
def execute_sync(self):
|
||||
"""Executes sub-operations sequentially in a synchronous blocking manner."""
|
||||
for op in self.sub_ops.values():
|
||||
for op in self.sub_ops:
|
||||
assert not op.async_mode
|
||||
op.call_sync(context=self.context)
|
||||
|
||||
|
|
|
|||
|
|
@ -20,7 +20,7 @@ class MCPTool(BaseOp):
|
|||
self,
|
||||
mcp_server: str = "",
|
||||
tool_name: str = "",
|
||||
enable_tool_response: bool = True,
|
||||
save_response_result: bool = True,
|
||||
parameter_required: List[str] | None = None,
|
||||
parameter_optional: List[str] | None = None,
|
||||
parameter_deleted: List[str] | None = None,
|
||||
|
|
@ -31,7 +31,7 @@ class MCPTool(BaseOp):
|
|||
):
|
||||
|
||||
super().__init__(
|
||||
enable_tool_response=enable_tool_response,
|
||||
save_response_result=save_response_result,
|
||||
max_retries=max_retries,
|
||||
raise_exception=raise_exception,
|
||||
**kwargs,
|
||||
|
|
|
|||
325
tests/test_op_composition.py
Normal file
325
tests/test_op_composition.py
Normal file
|
|
@ -0,0 +1,325 @@
|
|||
"""
|
||||
Unit tests for BaseOp and operator composition (>>, <<, |).
|
||||
Tests asynchronous execution mode.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
|
||||
from reme_ai.core.op import BaseOp
|
||||
from reme_ai.core.schema import ToolCall, ToolAttr
|
||||
|
||||
|
||||
class AddOp(BaseOp):
|
||||
"""Simple operator that adds a value to a number in context."""
|
||||
|
||||
def __init__(self, value: int = 1, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.value = value
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
return ToolCall(
|
||||
**{
|
||||
"name": self.name,
|
||||
"description": f"Add {self.value} to input",
|
||||
"parameters": ToolAttr(
|
||||
**{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"number": {"type": "integer", "description": "Input number"},
|
||||
},
|
||||
"required": ["number"],
|
||||
},
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self):
|
||||
"""Async execution: add value to input number."""
|
||||
self.context["number"] += self.value
|
||||
self.output = self.context["number"]
|
||||
|
||||
|
||||
class MultiplyOp(BaseOp):
|
||||
"""Simple operator that multiplies a number in context."""
|
||||
|
||||
def __init__(self, factor: int = 2, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.factor = factor
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
return ToolCall(
|
||||
**{
|
||||
"name": self.name,
|
||||
"description": f"Multiply by {self.factor}",
|
||||
"parameters": ToolAttr(
|
||||
**{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"number": {"type": "integer", "description": "Input number"},
|
||||
},
|
||||
"required": ["number"],
|
||||
},
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self):
|
||||
"""Async execution: multiply input number."""
|
||||
self.context["number"] *= self.factor
|
||||
self.output = self.context["number"]
|
||||
|
||||
|
||||
class AppendOp(BaseOp):
|
||||
"""Operator that appends a value to a list in context."""
|
||||
|
||||
def __init__(self, value: str = "", **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.value = value
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
return ToolCall(
|
||||
**{
|
||||
"name": self.name,
|
||||
"description": f"Append {self.value} to list",
|
||||
"parameters": ToolAttr(
|
||||
**{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"items": {"type": "array", "description": "List of items"},
|
||||
},
|
||||
"required": ["items"],
|
||||
},
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self):
|
||||
"""Async execution: append value to list."""
|
||||
self.context["items"].append(self.value)
|
||||
self.output = self.context["items"]
|
||||
|
||||
|
||||
async def test_basic_async_call():
|
||||
"""Test basic asynchronous operator execution."""
|
||||
op = AddOp(value=5, name="add_5")
|
||||
await op.call(number=10)
|
||||
number = op.context["number"]
|
||||
assert number == 15, f"Expected context result 15, got {number}"
|
||||
print("✓ test_basic_async_call passed")
|
||||
|
||||
|
||||
async def test_sequential_composition_async():
|
||||
"""Test >> operator for sequential composition in async mode."""
|
||||
add_op = AddOp(value=5, name="add_5")
|
||||
multiply_op = MultiplyOp(factor=2, name="multiply_2")
|
||||
composed = add_op >> multiply_op
|
||||
await composed.call(number=10)
|
||||
|
||||
# (10 + 5) * 2 = 30
|
||||
assert composed.context["number"] == 30, f"Expected 30, got {composed.context['number']}"
|
||||
print("✓ test_sequential_composition_async passed")
|
||||
|
||||
|
||||
async def test_parallel_composition_async():
|
||||
"""Test | operator for parallel composition in async mode."""
|
||||
append_a = AppendOp(value="A", name="append_a")
|
||||
append_b = AppendOp(value="B", name="append_b")
|
||||
append_c = AppendOp(value="C", name="append_c")
|
||||
|
||||
composed = append_a | append_b | append_c
|
||||
|
||||
await composed.call(items=[])
|
||||
|
||||
# All should append to the list
|
||||
items = composed.context["items"]
|
||||
assert len(items) == 3, f"Expected 3 items, got {len(items)}"
|
||||
assert set(items) == {"A", "B", "C"}, f"Expected A,B,C, got {items}"
|
||||
print("✓ test_parallel_composition_async passed")
|
||||
|
||||
|
||||
async def test_add_sub_ops_async():
|
||||
"""Test << operator for adding sub-operations in async mode."""
|
||||
parent_op = BaseOp(name="parent")
|
||||
child1 = AddOp(value=5, name="child1")
|
||||
child2 = MultiplyOp(factor=2, name="child2")
|
||||
|
||||
_ = parent_op << child1
|
||||
_ = parent_op << child2
|
||||
|
||||
assert len(parent_op.sub_ops) == 2, f"Expected 2 sub_ops, got {len(parent_op.sub_ops)}"
|
||||
sub_op_names = [op.name for op in parent_op.sub_ops]
|
||||
assert "child1" in sub_op_names, "child1 not in sub_ops"
|
||||
assert "child2" in sub_op_names, "child2 not in sub_ops"
|
||||
print("✓ test_add_sub_ops_async passed")
|
||||
|
||||
|
||||
async def test_add_sub_ops_dict():
|
||||
"""Test << operator with dictionary of operations."""
|
||||
parent_op = BaseOp(name="parent")
|
||||
ops_dict = {
|
||||
"add": AddOp(value=5, name="add"),
|
||||
"multiply": MultiplyOp(factor=2, name="multiply"),
|
||||
}
|
||||
|
||||
_ = parent_op << ops_dict
|
||||
|
||||
assert len(parent_op.sub_ops) == 2, f"Expected 2 ops_dict, got {len(parent_op.sub_ops)}"
|
||||
sub_op_names = [op.name for op in parent_op.sub_ops]
|
||||
assert "add" in sub_op_names, "add not in ops_dict"
|
||||
assert "multiply" in sub_op_names, "multiply not in ops_dict"
|
||||
print("✓ test_add_sub_ops_dict passed")
|
||||
|
||||
|
||||
async def test_add_sub_ops_list():
|
||||
"""Test << operator with list of operations."""
|
||||
parent_op = BaseOp(name="parent")
|
||||
sub_ops = [
|
||||
AddOp(value=5, name="add"),
|
||||
MultiplyOp(factor=2, name="multiply"),
|
||||
]
|
||||
|
||||
_ = parent_op << sub_ops
|
||||
|
||||
assert len(parent_op.sub_ops) == 2, f"Expected 2 sub_ops, got {len(parent_op.sub_ops)}"
|
||||
sub_op_names = [op.name for op in parent_op.sub_ops]
|
||||
assert "add" in sub_op_names, "add not in sub_ops"
|
||||
assert "multiply" in sub_op_names, "multiply not in sub_ops"
|
||||
print("✓ test_add_sub_ops_list passed")
|
||||
|
||||
|
||||
async def test_mixed_composition_async():
|
||||
"""Test mixing >> and | operators in async mode."""
|
||||
# (add_5 >> multiply_2) | (add_10 >> multiply_3)
|
||||
seq1 = AddOp(value=5, name="add_5") >> MultiplyOp(factor=2, name="multiply_2")
|
||||
seq2 = AddOp(value=10, name="add_10") >> MultiplyOp(factor=3, name="multiply_3")
|
||||
|
||||
composed = seq1 | seq2
|
||||
|
||||
await composed.call(number=10)
|
||||
|
||||
# Both sequences execute in parallel with shared context
|
||||
# seq1: (10 + 5) * 2 = 30
|
||||
# seq2: (30 + 10) * 3 = 120 (builds on seq1's result due to shared context)
|
||||
# The exact result depends on execution order and timing
|
||||
# With current implementation, result is 120
|
||||
assert composed.context["number"] == 120, f"Expected 120, got {composed.context['number']}"
|
||||
print("✓ test_mixed_composition_async passed")
|
||||
|
||||
|
||||
async def test_op_copy():
|
||||
"""Test operator copy functionality."""
|
||||
original = AddOp(value=5, name="original")
|
||||
copy_op = original.copy(name="copy")
|
||||
|
||||
assert copy_op.name == "copy", f"Expected name 'copy', got {copy_op.name}"
|
||||
assert copy_op.value == 5, f"Expected value 5, got {copy_op.value}"
|
||||
assert copy_op is not original, "Copy should be a different object"
|
||||
print("✓ test_op_copy passed")
|
||||
|
||||
|
||||
async def test_input_mapping():
|
||||
"""Test input_mapping parameter."""
|
||||
op = AddOp(
|
||||
value=5,
|
||||
name="add_5",
|
||||
input_mapping={"x": "number"}, # Map x to number
|
||||
)
|
||||
|
||||
await op.call(x=10) # Input is 'x' not 'number'
|
||||
|
||||
assert op.context["number"] == 15, f"Expected number=15, got {op.context['number']}"
|
||||
print("✓ test_input_mapping passed")
|
||||
|
||||
|
||||
async def test_output_mapping():
|
||||
"""Test output_mapping parameter."""
|
||||
op = AddOp(
|
||||
value=5,
|
||||
name="add_5",
|
||||
output_mapping={"number": "final_result"}, # Map number to final_result
|
||||
)
|
||||
|
||||
await op.call(number=10)
|
||||
|
||||
assert op.context["number"] == 15, f"Expected number=15, got {op.context['number']}"
|
||||
assert op.context["final_result"] == 15, f"Expected final_result=15, got {op.context['final_result']}"
|
||||
print("✓ test_output_mapping passed")
|
||||
|
||||
|
||||
async def test_validation_missing_required():
|
||||
"""Test that missing required inputs raise an error."""
|
||||
op = AddOp(value=5, name="add_5", raise_exception=True)
|
||||
|
||||
try:
|
||||
await op.call() # Missing 'number' field
|
||||
assert False, "Should have raised ValueError for missing required input"
|
||||
except ValueError as e:
|
||||
assert "number" in str(e), f"Expected error about 'number', got: {e}"
|
||||
print("✓ test_validation_missing_required passed")
|
||||
|
||||
|
||||
async def test_max_retries():
|
||||
"""Test max_retries parameter with failing operation."""
|
||||
|
||||
class FailingOp(BaseOp):
|
||||
"""An operation that always fails."""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.attempt_count = 0
|
||||
|
||||
def _build_tool_call(self) -> ToolCall:
|
||||
return ToolCall(
|
||||
**{
|
||||
"name": self.name,
|
||||
"description": "Always fails",
|
||||
"parameters": ToolAttr(**{"type": "object", "properties": {}}),
|
||||
"output": ToolAttr(
|
||||
**{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"result": ToolAttr(**{"type": "string", "description": "Result"}),
|
||||
},
|
||||
},
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self):
|
||||
self.attempt_count += 1
|
||||
raise RuntimeError(f"Attempt {self.attempt_count} failed")
|
||||
|
||||
op = FailingOp(max_retries=3, name="failing")
|
||||
|
||||
await op.call()
|
||||
|
||||
assert op.attempt_count == 3, f"Expected 3 attempts, got {op.attempt_count}"
|
||||
print("✓ test_max_retries passed")
|
||||
|
||||
|
||||
async def async_main():
|
||||
"""Run all async tests."""
|
||||
await test_basic_async_call()
|
||||
await test_sequential_composition_async()
|
||||
await test_parallel_composition_async()
|
||||
await test_add_sub_ops_async()
|
||||
await test_add_sub_ops_dict()
|
||||
await test_add_sub_ops_list()
|
||||
await test_mixed_composition_async()
|
||||
await test_op_copy()
|
||||
await test_input_mapping()
|
||||
await test_output_mapping()
|
||||
await test_validation_missing_required()
|
||||
await test_max_retries()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
print("Running BaseOp composition tests...\n")
|
||||
|
||||
# Async tests
|
||||
print("=== Asynchronous Tests ===")
|
||||
asyncio.run(async_main())
|
||||
|
||||
print("\n" + "=" * 50)
|
||||
print("All tests passed! ✓")
|
||||
print("=" * 50)
|
||||
Loading…
Add table
Reference in a new issue