diff --git a/reme_ai/core/context/runtime_context.py b/reme_ai/core/context/runtime_context.py index b2362161..d7112e1c 100644 --- a/reme_ai/core/context/runtime_context.py +++ b/reme_ai/core/context/runtime_context.py @@ -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.""" diff --git a/reme_ai/core/op/base_op.py b/reme_ai/core/op/base_op.py index e5f8450e..26f1e7f5 100644 --- a/reme_ai/core/op/base_op.py +++ b/reme_ai/core/op/base_op.py @@ -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) diff --git a/reme_ai/core/op/parallel_op.py b/reme_ai/core/op/parallel_op.py index 746485ca..18b84d0c 100644 --- a/reme_ai/core/op/parallel_op.py +++ b/reme_ai/core/op/parallel_op.py @@ -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() diff --git a/reme_ai/core/op/sequential_op.py b/reme_ai/core/op/sequential_op.py index 79243cff..3dabb0c9 100644 --- a/reme_ai/core/op/sequential_op.py +++ b/reme_ai/core/op/sequential_op.py @@ -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) diff --git a/reme_ai/core/tool/mcp_tool.py b/reme_ai/core/tool/mcp_tool.py index f9afa029..90f6ad7b 100644 --- a/reme_ai/core/tool/mcp_tool.py +++ b/reme_ai/core/tool/mcp_tool.py @@ -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, diff --git a/tests/test_op_composition.py b/tests/test_op_composition.py new file mode 100644 index 00000000..8d32b51a --- /dev/null +++ b/tests/test_op_composition.py @@ -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)