feat(core): add agent and tool modules with search and execution capabilities

This commit is contained in:
jinli.yl 2026-01-05 20:14:44 +08:00
parent 67f39db57a
commit f270e2a099
27 changed files with 1019 additions and 60 deletions

View file

@ -0,0 +1,19 @@
"""Core module for ReMe AI framework."""
# pylint: disable=wrong-import-position
# flake8: noqa: F401
from . import agent
from . import config
from . import context
from . import embedding
from . import enumeration
from . import flow
from . import llm
from . import op
from . import schema
from . import service
from . import token_counter
from . import tool
from . import utils
from . import vector_store

View file

@ -0,0 +1,9 @@
"""Agent module providing chat operations."""
from .simple_chat import SimpleChat
from .stream_chat import StreamChat
__all__ = [
"StreamChat",
"SimpleChat",
]

View file

@ -0,0 +1,62 @@
"""Simple chat agent for non-streaming conversations."""
from loguru import logger
from ..context import C
from ..enumeration import Role
from ..op import BaseOp
from ..schema import Message, ToolCall
@C.register_op()
class SimpleChat(BaseOp):
"""Simple chat agent that handles non-streaming conversations."""
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": "simple chat agent",
"parameters": {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "query",
},
"messages": {
"type": "array",
"items": {
"type": "object",
"properties": {
"role": {
"type": "string",
"description": "role",
},
"content": {
"type": "string",
"description": "content",
},
},
"required": ["role", "content"],
},
},
},
"required": [],
},
},
)
async def execute(self):
if "query" in self.context:
messages = [
Message(role=Role.SYSTEM, content="You are a helpful assistant."),
Message(role=Role.USER, content=self.context.query),
]
elif "messages" in self.context:
messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages if m]
else:
raise ValueError("query or messages must be provided!")
logger.info(f"messages={messages}")
assistant_message = await self.llm.chat(messages=messages)
logger.info(f"assistant_message={assistant_message.simple_dump()}")
self.output = assistant_message.content

View file

@ -0,0 +1,64 @@
"""Streaming chat agent for real-time conversation streaming."""
from loguru import logger
from ..context import C
from ..enumeration import Role, ChunkEnum
from ..op import BaseOp
from ..schema import Message, ToolCall
@C.register_op()
class StreamChat(BaseOp):
"""Streaming chat agent that handles real-time conversation streaming."""
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": "simple chat agent",
"parameters": {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "query",
},
"messages": {
"type": "array",
"items": {
"type": "object",
"properties": {
"role": {
"type": "string",
"description": "role",
},
"content": {
"type": "string",
"description": "content",
},
},
"required": ["role", "content"],
},
},
},
"required": [],
},
},
)
async def execute(self):
"""Execute streaming chat operation with query or messages."""
if "query" in self.context:
messages = [
Message(role=Role.SYSTEM, content="You are a helpful assistant."),
Message(role=Role.USER, content=self.context.query),
]
elif "messages" in self.context:
messages = [Message(**m) if isinstance(m, dict) else m for m in self.context.messages if m]
else:
raise ValueError("query or messages must be provided!")
logger.info(f"messages={messages}")
async for stream_chunk in self.llm.stream_chat(messages):
if stream_chunk.chunk_type in [ChunkEnum.ANSWER, ChunkEnum.THINK, ChunkEnum.ERROR, ChunkEnum.TOOL]:
await self.context.add_stream_chunk(stream_chunk)

View file

@ -6,7 +6,7 @@ import os
from .context import C
from .flow import BaseFlow
from .schema import ServiceConfig, Response
from .utils import PydanticConfigParser, init_logger, execute_stream_task, run_coro_safely
from .utils import PydanticConfigParser, init_logger, execute_stream_task, run_coro_safely, load_env
class Application:
@ -52,6 +52,8 @@ class Application:
token_counter: Token counter configuration dictionary
**kwargs: Additional configuration arguments
"""
load_env()
self._update_env("REME_LLM_API_KEY", llm_api_key)
self._update_env("REME_LLM_BASE_URL", llm_api_base)
self._update_env("REME_EMBEDDING_API_KEY", embedding_api_key)
@ -195,9 +197,9 @@ class Application:
task = asyncio.create_task(flow.call(stream_queue=stream_queue, **kwargs))
async for chunk in execute_stream_task(
queue=stream_queue,
stream_queue=stream_queue,
task=task,
flow_name=name,
task_name=name,
as_bytes=False,
):
yield chunk

View file

@ -16,6 +16,10 @@ llm:
default:
backend: openai
model_name: qwen3-30b-a3b-instruct-2507
qwen3_max_instruct:
backend: openai
model_name: qwen3-max
temperature: 0.6
embedding_model:
@ -32,3 +36,8 @@ vector_store:
token_counter:
default:
backend: base
hf:
backend: hf
model_name: Qwen/Qwen3-Coder-30B-A3B-Instruct
use_mirror: true

48
reme_ai/core/main.py Normal file
View file

@ -0,0 +1,48 @@
"""ReMe application classes for simplified configuration and execution."""
import sys
from .application import Application
from .config import ReMeConfigParser
class ReMeApp(Application):
"""ReMe application with config file support and flow execution methods."""
def __init__(
self,
*args,
llm_api_key: str | None = None,
llm_api_base: str | None = None,
embedding_api_key: str | None = None,
embedding_api_base: str | None = None,
config_path: str | None = None,
enable_logo: bool = True,
**kwargs,
):
super().__init__(
*args,
llm_api_key=llm_api_key,
llm_api_base=llm_api_base,
embedding_api_key=embedding_api_key,
embedding_api_base=embedding_api_base,
service_config=None,
parser=ReMeConfigParser,
config_path=config_path,
enable_logo=enable_logo,
**kwargs,
)
async def async_execute(self, name: str, **kwargs) -> dict:
"""Execute a flow asynchronously and return the result as a dictionary."""
return (await self.execute_flow(name=name, **kwargs)).model_dump()
def main():
"""Main entry point for running ReMe application from command line."""
with ReMeApp(*sys.argv[1:]) as app:
app.run_service()
if __name__ == "__main__":
main()

View file

@ -40,6 +40,7 @@ class BaseOp:
token_counter: str | BaseTokenCounter = "default",
enable_cache: bool = False,
cache_path: str = "cache/op",
cache_expire_hours: float = 0.1,
sub_ops: dict[str, "BaseOp"] | list["BaseOp"] | Optional["BaseOp"] = None,
input_mapping: dict[str, str] | None = None,
output_mapping: dict[str, str] | None = None,
@ -62,6 +63,7 @@ class BaseOp:
self.enable_cache = enable_cache
self.cache_path = cache_path
self.cache_expire_hours = cache_expire_hours
self.sub_ops: list[BaseOp] = []
self.add_sub_ops(sub_ops)
@ -156,7 +158,10 @@ class BaseOp:
return None
keys = list(output_properties.keys())
return self.context[keys[0]]
if len(keys) >= 1 and keys[0] in self.context:
return self.context[keys[0]]
else:
return None
@output.setter
def output(self, value: Any):

View file

@ -1,6 +1,4 @@
"""ReMe application classes for simplified configuration and execution."""
import sys
"""ReMe classes for simplified configuration and execution."""
from .application import Application
from .config import ReMeConfigParser
@ -48,45 +46,3 @@ class ReMe(Application):
async def retrieve(self):
"""Execute retrieve operations."""
class ReMeApp(Application):
"""ReMe application with config file support and flow execution methods."""
def __init__(
self,
*args,
llm_api_key: str | None = None,
llm_api_base: str | None = None,
embedding_api_key: str | None = None,
embedding_api_base: str | None = None,
config_path: str | None = None,
enable_logo: bool = True,
**kwargs,
):
super().__init__(
*args,
llm_api_key=llm_api_key,
llm_api_base=llm_api_base,
embedding_api_key=embedding_api_key,
embedding_api_base=embedding_api_base,
service_config=None,
parser=ReMeConfigParser,
config_path=config_path,
enable_logo=enable_logo,
**kwargs,
)
async def async_execute(self, name: str, **kwargs) -> dict:
"""Execute a flow asynchronously and return the result as a dictionary."""
return (await self.execute_flow(name=name, **kwargs)).model_dump()
def main():
"""Main entry point for running ReMe application from command line."""
with ReMeApp(*sys.argv[1:]) as app:
app.run_service()
if __name__ == "__main__":
main()

View file

@ -51,15 +51,14 @@ class HttpService(BaseService):
tool_call, request_model = self._prepare_route(flow)
async def execute_stream_endpoint(request: request_model) -> StreamingResponse:
queue = asyncio.Queue()
# Start flow as a background task
task = asyncio.create_task(flow.call(stream_queue=queue, **request.model_dump(exclude_none=True)))
stream_queue = asyncio.Queue()
task = asyncio.create_task(flow.call(stream_queue=stream_queue, **request.model_dump(exclude_none=True)))
async def generate_stream() -> AsyncGenerator[bytes, None]:
async for chunk in execute_stream_task(
queue=queue,
stream_queue=stream_queue,
task=task,
flow_name=tool_call.name,
task_name=tool_call.name,
as_bytes=True,
):
yield chunk

View file

@ -1,7 +1,9 @@
"""tool"""
from .mcp_tool import MCPTool
from . import search
__all__ = [
"MCPTool",
"search",
]

View file

@ -0,0 +1,9 @@
"""execute tool"""
from .execute_code import ExecuteCode
from .execute_shell import ExecuteShell
__all__ = [
"ExecuteCode",
"ExecuteShell",
]

View file

@ -0,0 +1,43 @@
"""Code execution tool for running Python code dynamically.
This module provides an operation that can execute Python code strings
and return the output or error messages.
"""
from ...context import C
from ...op import BaseOp
from ...schema import ToolCall
from ...utils import exec_code
@C.register_op()
class ExecuteCode(BaseOp):
"""Operation for executing Python code dynamically.
This operation takes Python code as input, executes it in a safe context,
and returns the output or any error messages that occur during execution.
"""
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": self.get_prompt("tool"),
"parameters": {
"type": "object",
"properties": {
"code": {
"type": "string",
"description": "code",
},
},
"required": ["code"],
},
},
)
async def execute(self):
self.execute_sync()
def execute_sync(self):
self.output = exec_code(self.context.code)

View file

@ -0,0 +1,5 @@
tool: |
Execute python code can be used in scenarios such as analysis or calculation, and the final result can be printed using the `print` function.
tool_zh: |
执行 Python 代码可用于分析或计算等场景,最终结果可以使用 print 函数输出。

View file

@ -0,0 +1,49 @@
"""Shell command execution tool.
This module provides an operation that can execute shell commands
asynchronously and return the output, error, and exit code.
"""
from ...context import C
from ...op import BaseOp
from ...schema import ToolCall
from ...utils import run_shell_command
@C.register_op()
class ExecuteShell(BaseOp):
"""Operation for executing shell commands asynchronously.
This operation takes a shell command as input, executes it asynchronously,
and returns the stdout, stderr, and exit code in a formatted result.
"""
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": self.get_prompt("tool"),
"parameters": {
"type": "object",
"properties": {
"command": {
"type": "string",
"description": "command",
},
},
"required": ["command"],
},
},
)
async def execute(self):
command: str = self.context.command
stdout, stderr, return_code = await run_shell_command(command)
result_parts = [
f"Command: {command}",
f"Output: {stdout if stdout else '(empty)'}",
f"Error: {stderr if stderr else '(none)'}",
f"Exit Code: {return_code if return_code is not None else '(none)'}",
]
self.output = "\n".join(result_parts)

View file

@ -0,0 +1,7 @@
tool: |
A tool capable of executing shell commands can use `pwd` to check the current location, `cd` to navigate to a new directory, `ls` to view the contents of a directory, and execute scripts.
Note that the starting directory is always the same each time the tool is invoked. If you need to perform multiple operations within a specific directory, you must include the full path in each command, for example: `cd aa/bb && bash xxx`.
tool_zh: |
一个能够执行 Shell 命令的工具可以使用 pwd 查看当前所在位置,使用 cd 切换到新目录,使用 ls 查看目录内容,并可执行脚本。
请注意,每次调用该工具时,起始目录始终相同。如果你需要在某个特定目录中执行多个操作,必须在每条命令中包含完整路径,例如:cd aa/bb && bash xxx。

View file

@ -0,0 +1,11 @@
"""search tool"""
from .dashscope_search import DashscopeSearch
from .mock_search import MockSearch
from .tavily_search import TavilySearch
__all__ = [
"DashscopeSearch",
"TavilySearch",
"MockSearch",
]

View file

@ -0,0 +1,111 @@
"""Dashscope web search tool.
This module provides an operation that uses Alibaba Cloud's Dashscope API
to perform web searches with various search strategies.
"""
import os
from typing import Literal
from loguru import logger
from ...context import C
from ...op import BaseOp
from ...schema import ToolCall
@C.register_op()
class DashscopeSearch(BaseOp):
"""Operation for performing web searches using Dashscope API.
This operation uses Alibaba Cloud's Dashscope service to search the web
with support for different search strategies (turbo, max, agent) and
optional role-based prompting.
"""
def __init__(
self,
model: str = "qwen-plus", # qwen-flash
search_strategy: Literal["turbo", "max", "agent"] = "turbo", # agent only for qwen3-max
enable_role_prompt: bool = True,
**kwargs,
):
super().__init__(**kwargs)
self.model: str = model
self.search_strategy: Literal["turbo", "max", "agent"] = search_strategy
self.enable_role_prompt: bool = enable_role_prompt
# see ref: https://help.aliyun.com/zh/model-studio/web-search
self.api_key = os.getenv("DASHSCOPE_API_KEY", "")
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": self.get_prompt("tool"),
"parameters": {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "query",
},
},
"required": ["query"],
},
},
)
async def execute(self):
query: str = self.context.query
if self.enable_cache:
cached_result = self.cache.load(query)
if cached_result:
self.output = cached_result["response_content"]
return
if self.enable_role_prompt:
user_query = self.prompt_format("role_prompt", query=query)
else:
user_query = query
logger.info(f"user_query={user_query}")
messages: list = [{"role": "user", "content": user_query}]
import dashscope
response = await dashscope.AioGeneration.call(
api_key=self.api_key,
model=self.model,
messages=messages,
enable_search=True,
search_options={
"forced_search": True,
"enable_source": True,
"enable_citation": False,
"search_strategy": self.search_strategy,
},
result_format="message",
)
search_results = []
response_content = ""
if response.output:
if response.output.search_info:
search_results = response.output.search_info.get("search_results", [])
if response.output.choices and len(response.output.choices) > 0:
response_content = response.output.choices[0].message.content
final_result = {
"query": query,
"search_results": search_results,
"response_content": response_content,
"model": self.model,
"search_strategy": self.search_strategy,
}
if self.enable_cache:
self.cache.save(query, final_result, expire_hours=self.cache_expire_hours)
self.output = final_result["response_content"]

View file

@ -0,0 +1,20 @@
tool: |
Use search keywords to retrieve relevant information from the internet.
If you have multiple keywords, please call this tool separately for each one.
tool_zh: |
使用搜索关键词从互联网检索相关信息。如果您有多个关键词,请分别为每个关键词单独调用此工具。
role_prompt: |
# user's question
{query}
# task
Extract the original content related to the user's query directly from the context, maintain accuracy, and avoid excessive processing.
role_prompt_zh: |
# 用户问题
{query}
# task
直接从上下文中提取与用户问题相关的原始内容,保持准确性,避免过度处理。

View file

@ -0,0 +1,64 @@
"""Mock search tool for testing purposes.
This module provides a mock search operation that generates simulated
search results using an LLM, useful for testing without making actual API calls.
"""
import json
import random
from loguru import logger
from ...context import C
from ...enumeration import Role
from ...op import BaseOp
from ...schema import ToolCall, Message
from ...utils import extract_content
@C.register_op()
class MockSearch(BaseOp):
"""Operation for generating mock search results.
This operation generates simulated search results using an LLM,
useful for testing and development without requiring actual search API access.
"""
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": self.get_prompt("tool"),
"parameters": {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "query",
},
},
"required": ["query"],
},
},
)
async def execute(self):
query: str = self.context.query
num_results: int = random.randint(0, 5)
messages = [
Message(
role=Role.SYSTEM,
content="You are a helpful assistant that generates realistic search results in JSON format.",
),
Message(
role=Role.USER,
content=self.prompt_format("mock_search_prompt", query=query, num_results=num_results),
),
]
logger.info(f"messages={messages}")
def callback_fn(message: Message):
return extract_content(message.content, "json")
search_results: str = await self.llm.chat(messages=messages, callback_fn=callback_fn)
self.output = json.dumps(search_results, ensure_ascii=False, indent=2)

View file

@ -0,0 +1,81 @@
tool: |
Use search keywords to retrieve relevant information from the internet.
If you have multiple keywords, please call this tool separately for each one.
tool_zh: |
使用搜索关键词从互联网检索相关信息。如果您有多个关键词,请分别为每个关键词单独调用此工具。
mock_search_prompt: |
# Task
Generate {num_results} realistic search results for the query: "{query}".
# Fields per item
Each result must be a JSON object with fields:
- snippet: 2-3 sentence summary
- title: page title
- url: realistic URL (e.g., https://example.com/article/title)
- hostname: domain (e.g., example.com)
- hostlogo: logo URL (e.g., https://example.com/logo.png) or empty string
# Requirements
- Ensure relevance to the query
- Use diverse, realistic sources
- Ensure well-formed URLs
- If no relevant results, return an empty array
# Output Format
First, think briefly about good sources and angles:
``` think
your brief reasoning here
```
Then output ONLY the JSON array wrapped in a json code block, nothing else:
``` json
[
{{
"snippet": "核心内容",
"title": "...",
"url": "...",
"hostname": "...",
"hostlogo": "..."
}}
]
```
mock_search_prompt_zh: |
# 任务
为查询“{query}”生成 {num_results} 条逼真的搜索结果。
# 每条结果的字段
每条结果必须是一个包含以下字段的 JSON 对象:
- snippet:2–3 句话的摘要
- title:网页标题
- url:逼真的 URL(例如:https://example.com/article/title)
- hostname:域名(例如:example.com)
- hostlogo:网站 logo 的 URL(例如:https://example.com/logo.png),若无则为空字符串
# 要求
- 确保结果与查询相关
- 使用多样且真实的来源
- 确保 URL 格式正确
- 若无相关结果,则返回空数组
# 输出格式
首先,简要思考合适的来源和角度:
``` think
你的简要推理写在这里
```
然后仅输出一个 JSON 数组,并用 json 代码块包裹,不要包含其他任何内容:
``` json
[
{
"snippet": "核心内容",
"title": "...",
"url": "...",
"hostname": "...",
"hostlogo": "..."
}
]
```

View file

@ -0,0 +1,119 @@
"""Tavily web search tool.
This module provides an operation that uses the Tavily API to perform
web searches and optionally extract content from search results.
"""
import json
import os
from loguru import logger
from ...context import C
from ...op import BaseOp
from ...schema import ToolCall
@C.register_op()
class TavilySearch(BaseOp):
"""Operation for performing web searches using Tavily API.
This operation uses the Tavily search service to find web content
and optionally extract raw content from the results, with configurable
character limits for individual items and total content.
"""
def __init__(
self,
enable_extract: bool = True,
item_max_char_count: int = 20000,
all_max_char_count: int = 50000,
**kwargs,
):
super().__init__(**kwargs)
self.enable_extract: bool = enable_extract
self.item_max_char_count: int = item_max_char_count
self.all_max_char_count: int = all_max_char_count
self._client = None
def _build_tool_call(self) -> ToolCall:
return ToolCall(
**{
"description": self.get_prompt("tool"),
"parameters": {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "query",
},
},
"required": ["query"],
},
},
)
@property
def client(self):
"""Get or create the Tavily async client instance.
Returns:
AsyncTavilyClient: The Tavily client instance, lazily initialized.
"""
if self._client is None:
from tavily import AsyncTavilyClient
self._client = AsyncTavilyClient(api_key=os.environ.get("TAVILY_API_KEY", ""))
return self._client
async def execute(self):
query: str = self.context.query
logger.info(f"tavily_search query={query}")
if self.enable_cache:
cached_result = self.cache.load(query)
if cached_result:
self.output = json.dumps(cached_result, ensure_ascii=False, indent=2)
return
response = await self.client.search(query=query)
logger.info(f"tavily_search response={response}")
if not self.enable_extract:
if not response.get("results"):
raise RuntimeError("tavily return empty result")
final_result = {item["url"]: item for item in response["results"]}
if self.enable_cache and final_result:
self.cache.save(query, final_result, expire_hours=self.cache_expire_hours)
self.output = json.dumps(final_result, ensure_ascii=False, indent=2)
return
url_info_dict = {item["url"]: item for item in response["results"]}
response_extract = await self.client.extract(urls=[item["url"] for item in response["results"]])
logger.info(f"tavily.response_extract: {response_extract}")
final_result = {}
all_char_count = 0
for item in response_extract["results"]:
url = item["url"]
raw_content: str = item["raw_content"]
if len(raw_content) > self.item_max_char_count:
raw_content = raw_content[: self.item_max_char_count]
if all_char_count + len(raw_content) > self.all_max_char_count:
raw_content = raw_content[: self.all_max_char_count - all_char_count]
if raw_content:
final_result[url] = url_info_dict[url]
final_result[url]["raw_content"] = raw_content
all_char_count += len(raw_content)
if not final_result:
raise RuntimeError("tavily return empty result")
if self.enable_cache and final_result:
self.cache.save(query, final_result, expire_hours=self.cache_expire_hours)
self.output = json.dumps(final_result, ensure_ascii=False, indent=2)

View file

@ -0,0 +1,6 @@
tool: |
Use search keywords to retrieve relevant information from the internet.
If you have multiple keywords, please call this tool separately for each one.
tool_zh: |
使用搜索关键词从互联网检索相关信息。如果您有多个关键词,请分别为每个关键词单独调用此工具。

View file

@ -4,6 +4,7 @@ from .cache_handler import CacheHandler
from .case_converter import snake_to_camel, camel_to_snake
from .common_utils import run_coro_safely, execute_stream_task
from .env_utils import load_env
from .execute_tuils import exec_code, run_shell_command
from .http_client import HttpClient
from .llm_utils import extract_content, format_messages
from .logger_utils import init_logger
@ -21,6 +22,8 @@ __all__ = [
"run_coro_safely",
"execute_stream_task",
"load_env",
"exec_code",
"run_shell_command",
"HttpClient",
"extract_content",
"format_messages",

View file

@ -26,9 +26,9 @@ def run_coro_safely(coro: Coroutine[Any, Any, Any]) -> Any | asyncio.Task[Any]:
async def execute_stream_task(
queue: asyncio.Queue,
stream_queue: asyncio.Queue,
task: asyncio.Task,
flow_name: str | None = None,
task_name: str | None = None,
as_bytes: bool = False,
) -> AsyncGenerator[str | bytes, None]:
"""
@ -38,9 +38,9 @@ async def execute_stream_task(
Properly manages errors and resource cleanup.
Args:
queue: Queue to receive StreamChunk objects from
stream_queue: Queue to receive StreamChunk objects from
task: Background task executing the flow
flow_name: Optional flow name for logging purposes
task_name: Optional flow name for logging purposes
as_bytes: If True, yield bytes for HTTP responses; if False, yield strings
Yields:
@ -51,7 +51,7 @@ async def execute_stream_task(
try:
while True:
# Wait for next chunk or check if task failed
get_chunk = asyncio.create_task(queue.get())
get_chunk = asyncio.create_task(stream_queue.get())
done, _ = await asyncio.wait({get_chunk, task}, return_when=asyncio.FIRST_COMPLETED)
if get_chunk in done:
@ -69,7 +69,7 @@ async def execute_stream_task(
break
except Exception as e:
log_msg = f"Stream error in {flow_name}: {e}" if flow_name else f"Stream error: {e}"
log_msg = f"Stream error in {task_name}: {e}" if task_name else f"Stream error: {e}"
logger.exception(log_msg)
err = StreamChunk(chunk_type=ChunkEnum.ERROR, chunk=str(e), done=True)

View file

@ -0,0 +1,60 @@
"""Utility functions for executing code and shell commands.
This module provides helper functions for running Python code and shell commands,
with support for async execution and output capture.
"""
import asyncio
import contextlib
from io import StringIO
async def run_shell_command(cmd: str, timeout: float | None = 30) -> tuple[str, str, int]:
"""Execute a shell command asynchronously.
Args:
cmd: The shell command to execute.
timeout: Maximum time to wait for command completion in seconds. None for no timeout.
Returns:
A tuple containing (stdout, stderr, return_code) as strings and integer.
"""
process = await asyncio.create_subprocess_shell(
cmd,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
)
if timeout:
stdout, stderr = await asyncio.wait_for(process.communicate(), timeout=timeout)
else:
stdout, stderr = await process.communicate()
return (
stdout.decode("utf-8", errors="ignore"),
stderr.decode("utf-8", errors="ignore"),
process.returncode,
)
def exec_code(code: str) -> str:
"""Execute Python code and capture the output.
Args:
code: The Python code string to execute.
Returns:
The captured stdout output, or the error message if execution fails.
"""
try:
redirected_output = StringIO()
with contextlib.redirect_stdout(redirected_output):
exec(code)
return redirected_output.getvalue()
except Exception as e:
return str(e)
except BaseException as e:
return str(e)

196
tests/test_tool.py Normal file
View file

@ -0,0 +1,196 @@
"""Tests for tool operations including search and execution tools.
This module contains test functions for various tool operations such as
search tools (Dashscope, Mock, Tavily) and execution tools (Code, Shell).
"""
# pylint: disable=too-many-statements
import asyncio
from reme_ai.core.reme import ReMe
ReMe()
def test_search():
"""Test search tool operations.
Tests DashscopeSearch, MockSearch, and TavilySearch operations
with a sample query to verify they work correctly.
"""
from reme_ai.core.tool.search import DashscopeSearch, MockSearch, TavilySearch
query = "今天杭州的天气如何?"
for op in [
DashscopeSearch(),
MockSearch(),
TavilySearch(),
]:
print("\n" + "=" * 60)
print(f"Testing {op.__class__.__name__}")
print("=" * 60)
print(f"Query: {query}")
asyncio.run(op.call(query=query))
print(f"Output:\n{op.output}")
def test_execute():
"""Test code and shell execution tool operations.
Tests ExecuteCode and ExecuteShell operations with various scenarios
including successful execution, syntax errors, runtime errors, and
invalid commands to verify error handling.
"""
from reme_ai.core.tool.execute import ExecuteCode, ExecuteShell
# Test ExecuteCode
print("\n" + "=" * 60)
print("Testing ExecuteCode")
print("=" * 60)
op = ExecuteCode()
code_to_execute = "print('hello world')"
print(f"Executing Python code: {code_to_execute}")
asyncio.run(op.call(code=code_to_execute))
print(f"Output:\n{op.output}")
# Test ExecuteCode with more complex code
print("\n" + "=" * 60)
print("Testing ExecuteCode with calculation")
print("=" * 60)
op = ExecuteCode()
code_to_execute = "result = sum(range(1, 11))\nprint(f'Sum of 1-10: {result}')"
print(f"Executing Python code:\n{code_to_execute}")
asyncio.run(op.call(code=code_to_execute))
print(f"Output:\n{op.output}")
# Test ExecuteShell
print("\n" + "=" * 60)
print("Testing ExecuteShell")
print("=" * 60)
op = ExecuteShell()
command = "ls"
print(f"Executing shell command: {command}")
asyncio.run(op.call(command=command))
print(f"Output:\n{op.output}")
# Test ExecuteShell with echo
print("\n" + "=" * 60)
print("Testing ExecuteShell with echo")
print("=" * 60)
op = ExecuteShell()
command = "echo 'Hello from shell!'"
print(f"Executing shell command: {command}")
asyncio.run(op.call(command=command))
print(f"Output:\n{op.output}")
# Test ExecuteCode with error (syntax error)
print("\n" + "=" * 60)
print("Testing ExecuteCode with syntax error (expected to fail)")
print("=" * 60)
op = ExecuteCode()
code_to_execute = "print('missing closing quote)"
print(f"Executing Python code with syntax error:\n{code_to_execute}")
asyncio.run(op.call(code=code_to_execute))
print(f"Output:\n{op.output}")
# Test ExecuteCode with runtime error
print("\n" + "=" * 60)
print("Testing ExecuteCode with runtime error (expected to fail)")
print("=" * 60)
op = ExecuteCode()
code_to_execute = "x = 1 / 0"
print(f"Executing Python code with runtime error:\n{code_to_execute}")
asyncio.run(op.call(code=code_to_execute))
print(f"Output:\n{op.output}")
# Test ExecuteCode with undefined variable
print("\n" + "=" * 60)
print("Testing ExecuteCode with undefined variable (expected to fail)")
print("=" * 60)
op = ExecuteCode()
code_to_execute = "print(undefined_variable)"
print(f"Executing Python code with undefined variable:\n{code_to_execute}")
asyncio.run(op.call(code=code_to_execute))
print(f"Output:\n{op.output}")
# Test ExecuteShell with invalid command
print("\n" + "=" * 60)
print("Testing ExecuteShell with invalid command (expected to fail)")
print("=" * 60)
op = ExecuteShell()
command = "this_command_does_not_exist"
print(f"Executing invalid shell command: {command}")
asyncio.run(op.call(command=command))
print(f"Output:\n{op.output}")
# Test ExecuteShell with command that returns non-zero exit code
print("\n" + "=" * 60)
print("Testing ExecuteShell with failing command (expected to fail)")
print("=" * 60)
op = ExecuteShell()
command = "ls /nonexistent_directory_12345"
print(f"Executing shell command that should fail: {command}")
asyncio.run(op.call(command=command))
print(f"Output:\n{op.output}")
print("\n" + "=" * 60)
print("All tests completed!")
print("=" * 60)
def test_simple_chat():
"""Test simple chat operation.
Tests the SimpleChat agent with a basic query to verify
it can process and respond to user input.
"""
from reme_ai.core.agent import SimpleChat
op = SimpleChat()
asyncio.run(op.call(query="你好"))
print(op.output)
async def test_stream_chat():
"""Test streaming chat operation.
Tests the StreamChat agent with a query to verify it can
process and stream responses in real-time using async operations.
"""
from reme_ai.core.agent import StreamChat
from reme_ai.core.utils import execute_stream_task
from reme_ai.core.context import RuntimeContext
from asyncio import Queue
op = StreamChat()
context = RuntimeContext(query="你好,详细介绍一下自己", stream_queue=Queue())
async def task():
await op.call(context)
await op.context.add_stream_done()
async for chunk in execute_stream_task(
stream_queue=context.stream_queue,
task=asyncio.create_task(task()),
task_name="test_stream_chat",
as_bytes=False,
):
print(chunk, end="")
if __name__ == "__main__":
# test_search()
# test_execute()
# test_simple_chat()
asyncio.run(test_stream_chat())