refactor(core): update BaseOp to improve sub-ops handling and tool call management

This commit is contained in:
jinli.yl 2026-01-01 13:10:21 +08:00
parent 472e069bc5
commit 566a773591
6 changed files with 395 additions and 48 deletions

View file

@ -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."""

View file

@ -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)

View file

@ -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()

View file

@ -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)

View file

@ -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,

View 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)