diff --git a/reme/components/agent_wrapper/as_agent_wrapper.py b/reme/components/agent_wrapper/as_agent_wrapper.py index 3ccc3df6..41c85ca1 100644 --- a/reme/components/agent_wrapper/as_agent_wrapper.py +++ b/reme/components/agent_wrapper/as_agent_wrapper.py @@ -42,6 +42,7 @@ from agentscope.tool import ( FunctionTool, Glob, Grep, + LocalBackend, Read, ToolBase, ToolChunk, @@ -66,6 +67,24 @@ _UUID_RE = re.compile( ) +class WorkspaceBackend(LocalBackend): + """LocalBackend whose reported cwd is the configured agent workspace. + + Some AgentScope builtin tools use ``backend.getcwd()`` for default search + paths or safety checks. Pinning it here keeps those operations aligned with + the cwd passed to Bash. Tools that require absolute file paths still keep + their own validation behavior. + """ + + def __init__(self, cwd: str) -> None: + super().__init__() + self._workspace_cwd = cwd + + async def getcwd(self) -> str: + """Return the configured workspace directory.""" + return self._workspace_cwd + + class BypassAnalysisBash(Bash): """Bash variant that delegates permission decisions to PermissionEngine. @@ -106,15 +125,51 @@ class AsAgentWrapper(BaseAgentWrapper): state = ToolResultState.SUCCESS if response.success else ToolResultState.ERROR return ToolChunk(content=[TextBlock(text=str(response.answer))], state=state) - tool = FunctionTool(func=run_job, name=job.name, description=job.description) + tool = FunctionTool(func=run_job, name=job.name, description=job.description, is_concurrency_safe=False) if job.parameters: tool.input_schema = job.parameters return tool - @classmethod - def _builtin_tools(cls) -> list[ToolBase]: - """Return built-in tools expected by local skills.""" - return [BypassAnalysisBash(), Edit(), Glob(), Grep(), Read(), Write()] + def _builtin_tools( + self, + names: list[str] | str | bool | None = "all", + *, + sequential_tool_calls: bool = False, + ) -> list[ToolBase]: + """Return selected AgentScope built-in tools rooted at ``self.cwd``.""" + cwd = str(self.cwd) + backend = WorkspaceBackend(cwd) + factories = { + "bash": lambda: BypassAnalysisBash(cwd=cwd, backend=backend), + "edit": lambda: Edit(backend=backend), + "glob": lambda: Glob(backend=backend), + "grep": lambda: Grep(backend=backend), + "read": lambda: Read(backend=backend), + "write": lambda: Write(backend=backend), + } + if names is False: + selected_names = [] + elif names is True or names is None or names == "all": + selected_names = list(factories) + elif names in ("none", "no", "false"): + selected_names = [] + elif isinstance(names, str): + selected_names = [names] + else: + selected_names = names + + tools: list[ToolBase] = [] + for name in selected_names: + key = name.lower() + if key not in factories: + allowed = ", ".join(factories) + raise ValueError(f"Unknown builtin tool {name!r}; expected one of: {allowed}") + tools.append(factories[key]()) + + if sequential_tool_calls: + for tool in tools: + tool.is_concurrency_safe = False + return tools @property def session_path(self) -> Path: @@ -220,8 +275,15 @@ class AsAgentWrapper(BaseAgentWrapper): resolved_jobs = self._resolve_job_tools(job_tools) skills = self._resolve_skills(kwargs.get("skills")) tool_context_id = kwargs.get("tool_context_id") + sequential_tool_calls = bool(kwargs.get("sequential_tool_calls", True)) + builtin_tools = kwargs.get("builtin_tools", "all") + if "builtin_tools" not in kwargs and not bool(kwargs.get("use_builtin_tools", True)): + builtin_tools = [] + tools: list[ToolBase] = [] + tools.extend(self._builtin_tools(builtin_tools, sequential_tool_calls=sequential_tool_calls)) + tools.extend(self._make_tool(job, tool_context_id) for job in resolved_jobs) toolkit = kwargs.get("toolkit") or Toolkit( - tools=[*self._builtin_tools(), *(self._make_tool(job, tool_context_id) for job in resolved_jobs)], + tools=tools, skills_or_loaders=skills, ) diff --git a/reme/components/agent_wrapper/base_agent_wrapper.py b/reme/components/agent_wrapper/base_agent_wrapper.py index f97559b3..fd4ec2ae 100644 --- a/reme/components/agent_wrapper/base_agent_wrapper.py +++ b/reme/components/agent_wrapper/base_agent_wrapper.py @@ -20,6 +20,23 @@ class BaseAgentWrapper(BaseComponent): component_type = ComponentEnum.AGENT_WRAPPER + def __init__(self, cwd: str | Path | None = None, **kwargs) -> None: + super().__init__(**kwargs) + self._cwd = cwd + + @property + def cwd(self) -> Path: + """Working directory shared by the agent's shell and file tools. + + Defaults to the project root (the workspace) — the same directory + Claude Code has always used. Override via the ``cwd`` init argument; + a relative value resolves against the workspace root. + """ + if not self._cwd: + return self.project_path + cwd = Path(self._cwd) + return cwd if cwd.is_absolute() else (self.workspace_path / cwd) + def set_system_prompt(self, prompt: str) -> "BaseAgentWrapper": """Set the agent's system prompt. Returns self for chaining.""" self.kwargs["system_prompt"] = prompt diff --git a/reme/components/agent_wrapper/cc_agent_wrapper.py b/reme/components/agent_wrapper/cc_agent_wrapper.py index e35ff886..f27bd5a0 100644 --- a/reme/components/agent_wrapper/cc_agent_wrapper.py +++ b/reme/components/agent_wrapper/cc_agent_wrapper.py @@ -277,7 +277,7 @@ class CcAgentWrapper(BaseAgentWrapper): ) opts.env.update(extra_env_dict) self.session_path.mkdir(parents=True, exist_ok=True) - opts.cwd = opts.cwd or self.project_path + opts.cwd = opts.cwd or self.cwd claude_config_dir = self.session_path / "claude_config" opts.env.setdefault("CLAUDE_CONFIG_DIR", str(claude_config_dir)) if opts.skills is not None: diff --git a/reme/components/service/__init__.py b/reme/components/service/__init__.py index 41d8f7b0..8ca82515 100644 --- a/reme/components/service/__init__.py +++ b/reme/components/service/__init__.py @@ -1,11 +1,13 @@ """Service components for exposing jobs via different protocols.""" from .base_service import BaseService +from .cli_service import CliService from .http_service import HttpService from .mcp_service import MCPService __all__ = [ "BaseService", + "CliService", "HttpService", "MCPService", ] diff --git a/reme/components/service/cli_service.py b/reme/components/service/cli_service.py new file mode 100644 index 00000000..e13cb1ca --- /dev/null +++ b/reme/components/service/cli_service.py @@ -0,0 +1,117 @@ +"""CLI service: run one configured job locally, then exit.""" + +import asyncio +import json +import sys +from typing import TYPE_CHECKING, Any + +from .base_service import BaseService +from ..component_registry import R +from ..job import BaseJob +from ...config import resolve_app_config +from ...schema import ApplicationConfig +from ...utils import get_logger + +if TYPE_CHECKING: + from ...application import Application + +_APP_CONFIG_KEYS = set(ApplicationConfig.model_fields) + + +def prepare_start_config(kwargs: dict) -> dict: + """Resolve ``reme start`` kwargs, translating top-level ``job=...`` into internal cli service config.""" + if "job" not in kwargs: + return resolve_app_config(**kwargs) + return _prepare_job_start_config(dict(kwargs)) + + +def should_precheck_start(config: dict) -> bool: + """CLI service does not bind a port, so it should skip service port prechecks.""" + service = config.get("service") + return not (isinstance(service, dict) and service.get("backend") == "cli") + + +def _prepare_job_start_config(kwargs: dict) -> dict: + """Translate ``reme start job=...`` args into the internal cli service fields.""" + job = kwargs.pop("job") + config_kwargs: dict = {} + job_args: dict = {} + + for key, value in kwargs.items(): + if key == "config" or key in _APP_CONFIG_KEYS: + config_kwargs[key] = value + else: + job_args[key] = value + + # One-shot CLI jobs should print only their answer by default. Reconfigure + # before resolve_app_config() so even config-loading logs stay off stdout. + if "log_to_console" not in config_kwargs: + get_logger(log_to_console=False, log_to_file=False, force_init=True) + + config = resolve_app_config(**config_kwargs) + if "enable_logo" not in config_kwargs: + config["enable_logo"] = False + if "log_to_console" not in config_kwargs: + config["log_to_console"] = False + service = dict(config.get("service") or {}) + service.update({"backend": "cli", "job": job, "job_args": job_args}) + config["service"] = service + return config + + +@R.register("cli") +class CliService(BaseService): + """Execute a single job through the normal application lifecycle without serving a port.""" + + def __init__( + self, + job: str = "", + job_args: dict[str, Any] | None = None, + show_metadata: bool = False, + **kwargs, + ): + super().__init__(**kwargs) + self.job = job + self.job_args = job_args or {} + self.show_metadata = show_metadata + + def build_service(self, app: "Application") -> None: + """No network framework is needed for local CLI execution.""" + self.service = None + + def add_job(self, job: BaseJob) -> bool: + """CLI execution does not register jobs; Application already owns them.""" + return False + + def start_service(self, app: "Application") -> None: + """Run the configured job once and print the same human-facing answer style as CLI clients.""" + asyncio.run(self._run_job(app)) + + def run_app(self, app: "Application") -> None: + """Bypass BaseService.add_jobs(), which is only meaningful for serving protocols.""" + self.build_service(app) + self.start_service(app) + + async def _run_job(self, app: "Application") -> None: + if not self.job: + raise ValueError("cli service requires service.job") + + await app.start() + try: + response = await app.run_job(self.job, **self.job_args) + output = self._format_response(response.answer, response.metadata) + if response.success: + print(output) + else: + print(output, file=sys.stderr) + raise SystemExit(1) + finally: + await app.close() + + def _format_response(self, answer: Any, metadata: dict | None) -> str: + if not isinstance(answer, str): + answer = json.dumps(answer, ensure_ascii=False, indent=2) + parts = [answer] + if self.show_metadata and metadata: + parts.append(json.dumps(metadata, ensure_ascii=False)) + return "\n".join(parts) diff --git a/reme/config/jinli_lme.yaml b/reme/config/jinli_lme.yaml index f842d44b..6b6f76d6 100644 --- a/reme/config/jinli_lme.yaml +++ b/reme/config/jinli_lme.yaml @@ -1,49 +1,23 @@ service: - backend: http + backend: cli + +workspace_dir: ${LME_WORKSPACE_DIR:-datasets/longmemeval/1} +session_dir: history_session +resource_dir: session +daily_dir: "" +digest_dir: "" jobs: update_index: backend: base - watch_dirs: [daily_dir, digest_dir, resource_dir] -# watch_dirs: [daily_dir, digest_dir] - watch_suffixes: [md, jsonl] -# watch_suffixes: [md] + watch_dirs: [resource_dir] + watch_suffixes: [md, json, jsonl] steps: + - backend: clear_store_step - backend: init_changes_step monitor_type: file_store monitor_name: default dispatch_steps: [update_index_step] - - backend: watch_changes_step - dispatch_steps: - - backend: update_index_step - persist: False - - auto_memory: - backend: base - description: "Auto-memory: record conversation facts into a daily note" - parameters: - type: object - properties: - messages: - type: array - description: "messages" - items: - type: object - session_id: - type: string - description: "source conversation session identifier" - default: "" - memory_hint: - type: string - description: "optional hint" - date: - type: string - description: "YYYY-MM-DD daily note date; empty = infer from message timestamps or today" - default: "" - required: - - messages - steps: - - backend: auto_memory_step version: backend: base @@ -63,14 +37,6 @@ jobs: query: type: string description: "search query" - limit: - type: integer - description: "max results" - default: 5 - min_score: - type: number - description: "min fused score" - default: 0.0 start_date: type: string description: "optional inclusive start date filter (YYYY-MM-DD); results earlier than this date are excluded" @@ -86,6 +52,46 @@ jobs: expand_links: true max_links_per_direction: 10 + add_draft: + backend: base + description: "Append text to the current draft list." + parameters: + type: object + properties: + text: + type: string + description: "draft text to append" + required: + - text + steps: + - backend: add_draft_step + + read_all_draft: + backend: base + description: "Read all draft text previously appended in the current tool context." + parameters: + type: object + properties: { } + steps: + - backend: read_all_draft_step + + python_execute: + backend: base + description: "Execute Python code and return printed stdout." + parameters: + type: object + properties: + code: + type: string + description: "Python code to execute. Print the final result to stdout." + timeout: + type: number + description: "Execution timeout in seconds; defaults to 60." + required: + - code + steps: + - backend: python_execute_step + components: tokenizer: default: @@ -109,16 +115,15 @@ components: as_llm: default: backend: ${LLM_BACKEND:-openai} - model: ${LLM_MODEL_NAME:-qwen3.7-plus} + model: ${LLM_MODEL_NAME:-qwen3.7-max} stream: true - context_size: 200000 + context_size: 1000000 max_retries: 3 credential: api_key: ${LLM_API_KEY:-} base_url: ${LLM_BASE_URL:-} parameters: max_tokens: 65536 - thinking_enable: false agent_wrapper: default: @@ -128,14 +133,36 @@ components: react_config: max_iters: 30 context_config: - trigger_ratio: 0.8 + trigger_ratio: 0.89 reserve_ratio: 0.1 tool_result_limit: 50000 model_config: - max_retries: 1 + max_retries: 3 + + agentic_search_agentwrapper: + backend: agentscope + as_llm: default + cwd: session + permission_mode: bypass + builtin_tools: false + job_tools: + - search + - add_draft + - read_all_draft + - python_execute + sequential_tool_calls: true + react_config: + max_iters: 100 + context_config: + trigger_ratio: 0.89 + reserve_ratio: 0.1 + tool_result_limit: 1000000 + model_config: + max_retries: 3 + claude_code: backend: claude_code - model: ${CLAUDE_CODE_MODEL_NAME:-glm-5.1} + model: ${CLAUDE_CODE_MODEL_NAME:-glm-5.2} api_key: ${CLAUDE_CODE_API_KEY:-} base_url: ${CLAUDE_CODE_BASE_URL:-https://dashscope.aliyuncs.com/apps/anthropic} permission_mode: bypassPermissions @@ -144,23 +171,14 @@ components: default: backend: local - file_catalog: - default: - backend: local - resource: - backend: local - digest: - backend: local - dream: - backend: local - file_chunker: markdown: backend: markdown supported_extensions: [ "md" ] default: backend: default - supported_extensions: [ "jsonl" ] + supported_extensions: [ "json", "jsonl" ] + chunk_byte_size: 100000 keyword_index: default: @@ -171,7 +189,7 @@ components: default: backend: local store_name: local - # embedding_store: default - embedding_store: "" + embedding_store: default +# embedding_store: "" keyword_index: default file_graph: default diff --git a/reme/reme.py b/reme/reme.py index 30f133a7..728ca437 100644 --- a/reme/reme.py +++ b/reme/reme.py @@ -5,6 +5,7 @@ import sys from .application import Application from .components import R +from .components.service.cli_service import prepare_start_config, should_precheck_start from .config import parse_args, resolve_app_config from .enumeration import ComponentEnum from .utils import cli_find_reme, load_env, precheck_start, running_service_config @@ -63,8 +64,8 @@ def main(): action, kwargs = parse_args(*sys.argv[1:]) if action == "start": load_env() - kwargs = resolve_app_config(**kwargs) - if not precheck_start(kwargs.get("service")): + kwargs = prepare_start_config(kwargs) + if should_precheck_start(kwargs) and not precheck_start(kwargs.get("service")): return ReMe(**kwargs).run_app() elif action == "find_reme": diff --git a/reme/steps/common/__init__.py b/reme/steps/common/__init__.py index 5d938939..bdc07d58 100644 --- a/reme/steps/common/__init__.py +++ b/reme/steps/common/__init__.py @@ -5,6 +5,7 @@ from .demo import DemoEchoStep1, DemoEchoStep2 from .health_check import HealthCheckStep from .help import HelpStep from .llm_demo import LLMDemoStep +from .python_execute import PythonExecuteStep from .stream_demo import StreamDemoStep1, StreamDemoStep2 from .stream_llm_demo import StreamLLMDemoStep from .version import VersionStep @@ -16,6 +17,7 @@ __all__ = [ "HealthCheckStep", "HelpStep", "LLMDemoStep", + "PythonExecuteStep", "StreamDemoStep1", "StreamDemoStep2", "StreamLLMDemoStep", diff --git a/reme/steps/common/python_execute.py b/reme/steps/common/python_execute.py new file mode 100644 index 00000000..b98b48db --- /dev/null +++ b/reme/steps/common/python_execute.py @@ -0,0 +1,100 @@ +"""Execute Python code and return printed stdout.""" + +import asyncio +import sys +from dataclasses import dataclass +from typing import Any + +from ..base_step import BaseStep +from ...components import R + +DEFAULT_TIMEOUT = 60.0 + + +@dataclass(frozen=True) +class _PythonResult: + stdout: str + stderr: str + returncode: int | None + timed_out: bool = False + + +@R.register("python_execute_step") +class PythonExecuteStep(BaseStep): + """Run Python code in a subprocess and return stdout as the response answer.""" + + async def execute(self): + assert self.context is not None + + code = self.context.get("code", "") + timeout, timeout_error = self._parse_timeout(self.context.get("timeout", DEFAULT_TIMEOUT)) + if not isinstance(code, str) or not code.strip(): + self.context.response.success = False + self.context.response.answer = "code is required" + return self.context.response + if timeout_error: + self.context.response.success = False + self.context.response.answer = timeout_error + return self.context.response + + result = await self._run_python(code, timeout) + if result.timed_out: + self.context.response.success = False + self.context.response.answer = f"Python execution timed out after {timeout:g}s" + self.context.response.metadata.update( + { + "returncode": result.returncode, + "stderr": result.stderr, + "timeout": timeout, + }, + ) + return self.context.response + + stdout = result.stdout or "" + stderr = result.stderr or "" + self.context.response.success = result.returncode == 0 + self.context.response.answer = stdout if stdout or result.returncode == 0 else stderr + self.context.response.metadata.update( + { + "returncode": result.returncode, + "stderr": stderr, + "timeout": timeout, + }, + ) + return self.context.response + + async def _run_python(self, code: str, timeout: float) -> _PythonResult: + process = await asyncio.create_subprocess_exec( + sys.executable, + "-c", + code, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + cwd=self.workspace_path, + ) + try: + stdout, stderr = await asyncio.wait_for(process.communicate(), timeout=timeout) + return _PythonResult( + stdout=stdout.decode(), + stderr=stderr.decode(), + returncode=process.returncode, + ) + except TimeoutError: + process.kill() + stdout, stderr = await process.communicate() + return _PythonResult( + stdout=stdout.decode(), + stderr=stderr.decode(), + returncode=process.returncode, + timed_out=True, + ) + + @staticmethod + def _parse_timeout(raw: Any) -> tuple[float, str]: + try: + timeout = float(raw) + except (TypeError, ValueError): + return DEFAULT_TIMEOUT, "timeout must be a positive number" + if timeout <= 0: + return DEFAULT_TIMEOUT, "timeout must be a positive number" + return timeout, "" diff --git a/reme/steps/index/__init__.py b/reme/steps/index/__init__.py index 7b903e3a..8cbf69bd 100644 --- a/reme/steps/index/__init__.py +++ b/reme/steps/index/__init__.py @@ -1,12 +1,15 @@ """Index steps.""" +from .bm25_search import Bm25SearchStep from .clear_store import ClearStoreStep +from .draft import AddDraftStep, ReadAllDraftStep from .log_changes import LogChangesStep from .node_search import NodeSearchStep from .init_changes import InitChangesStep from .search import SearchStep from .traverse import TraverseStep from .update_changes import ChangeApplyStep, UpdateCatalogStep, UpdateIndexStep +from .vector_search import VectorSearchStep from .watch_changes import ( DEFAULT_LOW_POWER_POLL_MS, DEFAULT_WATCH_DEBOUNCE_MS, @@ -15,6 +18,8 @@ from .watch_changes import ( ) __all__ = [ + "AddDraftStep", + "Bm25SearchStep", "ChangeApplyStep", "ClearStoreStep", "DEFAULT_LOW_POWER_POLL_MS", @@ -23,9 +28,11 @@ __all__ = [ "InitChangesStep", "LogChangesStep", "NodeSearchStep", + "ReadAllDraftStep", "SearchStep", "TraverseStep", "UpdateCatalogStep", "UpdateIndexStep", + "VectorSearchStep", "WatchChangesStep", ] diff --git a/reme/steps/index/bm25_search.py b/reme/steps/index/bm25_search.py new file mode 100644 index 00000000..b9fc47f3 --- /dev/null +++ b/reme/steps/index/bm25_search.py @@ -0,0 +1,71 @@ +"""``bm25_search_step`` — plain BM25 keyword search with tool_context dedup.""" + +import datetime +from typing import Final + +from ..base_step import BaseStep +from ...components import R +from ...schema import FileChunk + +_MAX_CANDIDATES: Final = 200 + + +@R.register("bm25_search_step") +class Bm25SearchStep(BaseStep): + """BM25-only search: retrieve, filter by min_score, dedup by tool_context, truncate.""" + + TOOL_CONTEXTS_KEY: Final[str] = "tool_contexts" + SEARCH_SEEN_KEY: Final[str] = "search_seen_chunk_ids" + + def __init__(self, *args, seen_ttl_hours: float = 24, **kwargs): + super().__init__(*args, **kwargs) + self.seen_ttl_hours = seen_ttl_hours + + def _tool_context_store(self, tool_context_id: str) -> dict: + """Return the mutable state bucket for a tool context.""" + if self.app_context is not None: + contexts = self.app_context.metadata.setdefault(self.TOOL_CONTEXTS_KEY, {}) + else: + contexts = self.kwargs.setdefault(self.TOOL_CONTEXTS_KEY, {}) + return contexts.setdefault(tool_context_id, {}) + + def _dedupe_tool_context(self, chunks: list[FileChunk], tool_context_id: str, limit: int) -> list[FileChunk]: + """Drop chunks already returned for this tool_context within the TTL window.""" + now = datetime.datetime.now().timestamp() + ttl = float(self.seen_ttl_hours) * 60 * 60 + store = self._tool_context_store(tool_context_id) + seen: dict = store.get(self.SEARCH_SEEN_KEY, {}) + seen = {cid: ts for cid, ts in seen.items() if now - float(ts) < ttl} + + returned = [c for c in chunks if c.id not in seen][:limit] + for c in returned: + seen[c.id] = now + store[self.SEARCH_SEEN_KEY] = seen + return returned + + async def execute(self): + assert self.context is not None + query: str = (self.context.get("query", "") or "").strip() + limit: int = int(self.context.get("limit") or 5) + tool_context_id: str = (self.context.get("tool_context_id", "") or "").strip() + + if not query: + self.context.response.success = False + self.context.response.answer = "Error: query cannot be empty" + return self.context.response + assert limit > 0, f"limit must be positive, got {limit}" + + candidates = min(_MAX_CANDIDATES, max(1, limit * 5)) + results = await self.file_store.keyword_search(query, candidates, {}) + self.logger.info(f"[{self.name}] query={query!r} candidates={candidates} hits={len(results)}") + + if tool_context_id: + results = self._dedupe_tool_context(results, tool_context_id, limit) + else: + results = results[:limit] + + self.context.response.answer = "\n\n".join(c.text for c in results) + self.context.response.metadata["results"] = [ + c.model_dump(exclude_none=True, exclude={"embedding"}) for c in results + ] + return self.context.response diff --git a/reme/steps/index/draft.py b/reme/steps/index/draft.py new file mode 100644 index 00000000..a33d3db2 --- /dev/null +++ b/reme/steps/index/draft.py @@ -0,0 +1,63 @@ +"""Draft accumulation steps scoped by agent tool context.""" + +from typing import Final + +from ..base_step import BaseStep +from ...components import R + + +@R.register("add_draft_step") +class AddDraftStep(BaseStep): + """Append one draft text item to the current tool context.""" + + TOOL_CONTEXTS_KEY: Final[str] = "tool_contexts" + DRAFTS_KEY: Final[str] = "drafts" + + def _tool_context_store(self, tool_context_id: str) -> dict: + if self.app_context is not None: + contexts = self.app_context.metadata.setdefault(self.TOOL_CONTEXTS_KEY, {}) + else: + contexts = self.kwargs.setdefault(self.TOOL_CONTEXTS_KEY, {}) + return contexts.setdefault(tool_context_id, {}) + + async def execute(self): + assert self.context is not None + text = self.context.get("text", "") + tool_context_id: str = (self.context.get("tool_context_id", "") or "").strip() + + if not tool_context_id: + self.context.response.success = False + self.context.response.answer = "Error: tool_context_id is required" + return self.context.response + if text is None or not isinstance(text, str): + self.context.response.success = False + self.context.response.answer = "Error: text must be a string" + return self.context.response + + store = self._tool_context_store(tool_context_id) + drafts = store.setdefault(self.DRAFTS_KEY, []) + drafts.append(text) + + self.context.response.answer = text + self.context.response.metadata["draft_count"] = len(drafts) + return self.context.response + + +@R.register("read_all_draft_step") +class ReadAllDraftStep(AddDraftStep): + """Read all draft text items for the current tool context.""" + + async def execute(self): + assert self.context is not None + tool_context_id: str = (self.context.get("tool_context_id", "") or "").strip() + + if not tool_context_id: + self.context.response.success = False + self.context.response.answer = "Error: tool_context_id is required" + return self.context.response + + store = self._tool_context_store(tool_context_id) + drafts = store.setdefault(self.DRAFTS_KEY, []) + self.context.response.answer = "\n".join(drafts) + self.context.response.metadata["draft_count"] = len(drafts) + return self.context.response diff --git a/reme/steps/index/search.py b/reme/steps/index/search.py index a22942d7..8f29267c 100644 --- a/reme/steps/index/search.py +++ b/reme/steps/index/search.py @@ -2,6 +2,8 @@ import asyncio import datetime +import os +from typing import Final from ..base_step import BaseStep from ..file_io import extract_daily_date @@ -9,26 +11,37 @@ from ...components import R from ...schema import FileChunk from ...utils import expand_links, render_expansion_lines -_RRF_K = 60 -_MAX_CANDIDATES = 200 +_RRF_K: Final = 60 +_MAX_CANDIDATES: Final = 200 +_DEFAULT_LIMIT_ENV: Final = "REME_SEARCH_LIMIT" +_DEFAULT_LIMIT: Final = 5 + + +def _default_limit() -> int: + value = os.getenv(_DEFAULT_LIMIT_ENV) + if value is None: + return _DEFAULT_LIMIT + try: + return int(value) + except ValueError: + return _DEFAULT_LIMIT @R.register("search_step") class SearchStep(BaseStep): """Hybrid search: run vector + keyword in parallel, fuse via RRF, filter, truncate.""" + TOOL_CONTEXTS_KEY: Final[str] = "tool_contexts" + SEARCH_SEEN_KEY: Final[str] = "search_seen_chunk_ids" + def __init__( self, *args, - tool_contexts_key: str = "tool_contexts", - search_seen_key: str = "search_seen_chunk_ids", - tool_context_chunk_ttl_hours: float = 24, + seen_ttl_hours: float = 24, **kwargs, ): super().__init__(*args, **kwargs) - self.tool_contexts_key = tool_contexts_key - self.search_seen_key = search_seen_key - self.tool_context_chunk_ttl_hours = tool_context_chunk_ttl_hours + self.seen_ttl_hours = seen_ttl_hours @staticmethod def _rrf_merge( @@ -81,9 +94,9 @@ class SearchStep(BaseStep): def _tool_context_store(self, tool_context_id: str) -> dict: """Return the mutable state bucket for a tool context.""" if self.app_context is not None: - contexts = self.app_context.metadata.setdefault(self.tool_contexts_key, {}) + contexts = self.app_context.metadata.setdefault(self.TOOL_CONTEXTS_KEY, {}) else: - contexts = self.kwargs.setdefault(self.tool_contexts_key, {}) + contexts = self.kwargs.setdefault(self.TOOL_CONTEXTS_KEY, {}) return contexts.setdefault(tool_context_id, {}) def _dedupe_tool_context( @@ -96,15 +109,15 @@ class SearchStep(BaseStep): ttl = float( self.kwargs.get( "tool_context_chunk_ttl_seconds", - float(self.tool_context_chunk_ttl_hours) * 60 * 60, + float(self.seen_ttl_hours) * 60 * 60, ), ) store = self._tool_context_store(tool_context_id) - seen = store.get(self.search_seen_key, {}) + seen = store.get(self.SEARCH_SEEN_KEY, {}) if not isinstance(seen, dict): seen = dict.fromkeys(seen, now) before_expire = len(seen) - store[self.search_seen_key] = seen = {chunk_id: ts for chunk_id, ts in seen.items() if now - float(ts) < ttl} + store[self.SEARCH_SEEN_KEY] = seen = {chunk_id: ts for chunk_id, ts in seen.items() if now - float(ts) < ttl} seen_before = len(seen) unvisited = [chunk for chunk in chunks if chunk.id not in seen] @@ -124,7 +137,7 @@ class SearchStep(BaseStep): async def execute(self): assert self.context is not None query: str = (self.context.get("query", "") or "").strip() - limit: int = int(self.context.get("limit") or 5) + limit: int = int(self.context.get("limit") or _default_limit()) min_score: float = float(self.context.get("min_score") or 0.0) vector_weight: float = float(self.kwargs.get("vector_weight", 0.7)) candidate_multiplier: float = float(self.kwargs.get("candidate_multiplier", 5.0)) diff --git a/reme/steps/index/vector_search.py b/reme/steps/index/vector_search.py new file mode 100644 index 00000000..673252c2 --- /dev/null +++ b/reme/steps/index/vector_search.py @@ -0,0 +1,71 @@ +"""``vector_search_step`` — plain vector search with tool_context dedup.""" + +import datetime +from typing import Final + +from ..base_step import BaseStep +from ...components import R +from ...schema import FileChunk + +_MAX_CANDIDATES: Final = 200 + + +@R.register("vector_search_step") +class VectorSearchStep(BaseStep): + """Vector-only search: retrieve, filter by min_score, dedup by tool_context, truncate.""" + + TOOL_CONTEXTS_KEY: Final[str] = "tool_contexts" + SEARCH_SEEN_KEY: Final[str] = "search_seen_chunk_ids" + + def __init__(self, *args, seen_ttl_hours: float = 24, **kwargs): + super().__init__(*args, **kwargs) + self.seen_ttl_hours = seen_ttl_hours + + def _tool_context_store(self, tool_context_id: str) -> dict: + """Return the mutable state bucket for a tool context.""" + if self.app_context is not None: + contexts = self.app_context.metadata.setdefault(self.TOOL_CONTEXTS_KEY, {}) + else: + contexts = self.kwargs.setdefault(self.TOOL_CONTEXTS_KEY, {}) + return contexts.setdefault(tool_context_id, {}) + + def _dedupe_tool_context(self, chunks: list[FileChunk], tool_context_id: str, limit: int) -> list[FileChunk]: + """Drop chunks already returned for this tool_context within the TTL window.""" + now = datetime.datetime.now().timestamp() + ttl = float(self.seen_ttl_hours) * 60 * 60 + store = self._tool_context_store(tool_context_id) + seen: dict = store.get(self.SEARCH_SEEN_KEY, {}) + seen = {cid: ts for cid, ts in seen.items() if now - float(ts) < ttl} + + returned = [c for c in chunks if c.id not in seen][:limit] + for c in returned: + seen[c.id] = now + store[self.SEARCH_SEEN_KEY] = seen + return returned + + async def execute(self): + assert self.context is not None + query: str = (self.context.get("query", "") or "").strip() + limit: int = int(self.context.get("limit") or 5) + tool_context_id: str = (self.context.get("tool_context_id", "") or "").strip() + + if not query: + self.context.response.success = False + self.context.response.answer = "Error: query cannot be empty" + return self.context.response + assert limit > 0, f"limit must be positive, got {limit}" + + candidates = min(_MAX_CANDIDATES, max(1, limit * 5)) + results = await self.file_store.vector_search(query, candidates, {}) + self.logger.info(f"[{self.name}] query={query!r} candidates={candidates} hits={len(results)}") + + if tool_context_id: + results = self._dedupe_tool_context(results, tool_context_id, limit) + else: + results = results[:limit] + + self.context.response.answer = "\n\n".join(c.text for c in results) + self.context.response.metadata["results"] = [ + c.model_dump(exclude_none=True, exclude={"embedding"}) for c in results + ] + return self.context.response diff --git a/tests/unit/test_common_steps.py b/tests/unit/test_common_steps.py index ee9bf979..d6c4a13f 100644 --- a/tests/unit/test_common_steps.py +++ b/tests/unit/test_common_steps.py @@ -8,11 +8,13 @@ import tempfile import warnings from reme.components.agent_wrapper import BaseAgentWrapper +from reme.components.application_context import ApplicationContext from reme.components.file_store import LocalFileStore from reme.schema import FileLink, FileNode from reme.steps.common.add import AddStep from reme.steps.common.health_check import _file_graph_status from reme.steps.common.llm_demo import LLMDemoStep +from reme.steps.common.python_execute import PythonExecuteStep from reme.steps.index import traverse as traverse_mod warnings.filterwarnings("ignore", category=DeprecationWarning, module="jieba") @@ -89,6 +91,51 @@ def test_add_step_rejects_invalid_inputs(): _run(run()) +def test_python_execute_step_returns_printed_stdout(tmp_path): + """python_execute returns stdout as answer and runs under workspace_dir.""" + + async def run(): + app_context = ApplicationContext(workspace_dir=str(tmp_path)) + step = PythonExecuteStep(app_context=app_context) + resp = await step(code="from pathlib import Path\nprint(Path.cwd().name)") + assert resp.success is True + assert resp.answer == f"{tmp_path.name}\n" + assert resp.metadata["returncode"] == 0 + assert resp.metadata["stderr"] == "" + print("✓ test_python_execute_step_returns_printed_stdout passed") + + _run(run()) + + +def test_python_execute_step_reports_stderr_on_failure(): + """python_execute captures traceback stderr instead of throwing.""" + + async def run(): + step = PythonExecuteStep() + resp = await step(code='raise RuntimeError("boom")') + assert resp.success is False + assert resp.metadata["returncode"] != 0 + assert "RuntimeError: boom" in resp.answer + assert "RuntimeError: boom" in resp.metadata["stderr"] + print("✓ test_python_execute_step_reports_stderr_on_failure passed") + + _run(run()) + + +def test_python_execute_step_times_out(): + """python_execute converts subprocess timeout into a failed response.""" + + async def run(): + step = PythonExecuteStep() + resp = await step(code="import time\ntime.sleep(1)", timeout=0.01) + assert resp.success is False + assert resp.answer == "Python execution timed out after 0.01s" + assert resp.metadata["timeout"] == 0.01 + print("✓ test_python_execute_step_times_out passed") + + _run(run()) + + class _FakeAgentWrapper(BaseAgentWrapper): """Capture reply kwargs without calling a real model.""" diff --git a/tests/unit/test_reme_cli.py b/tests/unit/test_reme_cli.py index 331ac8fc..7becea3a 100644 --- a/tests/unit/test_reme_cli.py +++ b/tests/unit/test_reme_cli.py @@ -1,8 +1,146 @@ """Tests for the ReMe CLI entry helpers.""" +import asyncio +from types import SimpleNamespace + +import pytest + +from reme.components.service import cli_service +from reme.components.service.cli_service import CliService from reme import reme as reme_module +def test_prepare_start_config_moves_unknown_start_args_to_job_args(monkeypatch): + """``reme start job=...`` is translated into a one-shot cli service config.""" + + monkeypatch.setattr( + cli_service, + "resolve_app_config", + lambda **kwargs: { + **kwargs, + "service": {"backend": "http", "host": "127.0.0.1"}, + }, + ) + + cfg = cli_service.prepare_start_config( + { + "config": "jinli_lme", + "workspace_dir": "/tmp/reme", + "job": "search", + "query": "hello", + "limit": 3, + }, + ) + + assert cfg["config"] == "jinli_lme" + assert cfg["workspace_dir"] == "/tmp/reme" + assert cfg["enable_logo"] is False + assert cfg["log_to_console"] is False + assert cfg["service"] == { + "backend": "cli", + "host": "127.0.0.1", + "job": "search", + "job_args": {"query": "hello", "limit": 3}, + } + + +def test_should_precheck_start_skips_cli_service(): + """CLI service is local execution and should not run port prechecks.""" + assert cli_service.should_precheck_start({"service": {"backend": "cli"}}) is False + assert cli_service.should_precheck_start({"service": {"backend": "http"}}) is True + + +def test_cli_service_runs_configured_job_and_closes_app(capsys): + """CLI service runs one local job through app lifecycle and prints its answer.""" + events = [] + + class FakeApp: + """Minimal app stub for exercising CliService lifecycle.""" + + async def start(self): + """Record app startup.""" + events.append("start") + + async def close(self): + """Record app shutdown.""" + events.append("close") + + async def run_job(self, name, **kwargs): + """Record job execution and return a successful response.""" + events.append(("run_job", name, kwargs)) + return SimpleNamespace(answer="found it", success=True, metadata={"hits": 1}) + + service = CliService(job="search", job_args={"query": "hello"}) + + service.start_service(FakeApp()) + + assert events == [ + "start", + ("run_job", "search", {"query": "hello"}), + "close", + ] + assert capsys.readouterr().out == "found it\n" + + +def test_cli_service_can_print_metadata_from_service_config(capsys): + """service.show_metadata controls optional CLI metadata output.""" + + class FakeApp: + """Minimal app stub for exercising metadata output.""" + + async def start(self): + """No-op app startup.""" + + async def close(self): + """No-op app shutdown.""" + + async def run_job(self, _name, **_kwargs): + """Return a successful response with metadata.""" + return SimpleNamespace(answer="found it", success=True, metadata={"hits": 1}) + + service = CliService(job="search", show_metadata=True) + + service.start_service(FakeApp()) + + assert capsys.readouterr().out == 'found it\n{"hits": 1}\n' + + +def test_cli_service_exits_nonzero_on_failed_response(capsys): + """Failed local CLI jobs write to stderr and produce a failing process status.""" + events = [] + + class FakeApp: + """Minimal app stub for exercising failure handling.""" + + async def start(self): + """Record app startup.""" + events.append("start") + + async def close(self): + """Record app shutdown.""" + events.append("close") + + async def run_job(self, name, **kwargs): + """Record job execution and return a failed response.""" + events.append(("run_job", name, kwargs)) + return SimpleNamespace(answer="boom", success=False, metadata={}) + + service = CliService(job="search", job_args={"query": "hello"}) + + with pytest.raises(SystemExit) as exc_info: + service.start_service(FakeApp()) + + assert exc_info.value.code == 1 + assert events == [ + "start", + ("run_job", "search", {"query": "hello"}), + "close", + ] + captured = capsys.readouterr() + assert captured.out == "" + assert captured.err == "boom\n" + + def test_call_server_passes_client_kwargs_to_client(monkeypatch, capsys): """CLI helper forwards connection options to the selected client.""" seen = {} @@ -37,8 +175,6 @@ def test_call_server_passes_client_kwargs_to_client(monkeypatch, capsys): query="hello", ) - import asyncio - asyncio.run(run()) assert seen["client_kwargs"] == {"host": "127.0.0.2", "port": 2444, "timeout": 1.5} @@ -74,8 +210,6 @@ def test_call_server_treats_show_metadata_as_client_kwarg(monkeypatch, capsys): async def run(): await reme_module.call_server("version", backend="http", show_metadata=True) - import asyncio - asyncio.run(run()) assert seen["client_kwargs"] == {"show_metadata": True} diff --git a/tests/unit/test_search_step.py b/tests/unit/test_search_step.py index 7d347255..57e65f6c 100644 --- a/tests/unit/test_search_step.py +++ b/tests/unit/test_search_step.py @@ -3,10 +3,11 @@ import asyncio from reme.components.file_store import BaseFileStore +from reme.components import ApplicationContext from reme.components.runtime_context import RuntimeContext from reme.enumeration import LinkScopeEnum from reme.schema import FileChunk, FileLink, FileNode -from reme.steps.index import SearchStep +from reme.steps.index import AddDraftStep, ReadAllDraftStep, SearchStep class FakeSearchStore(BaseFileStore): @@ -109,6 +110,26 @@ def test_search_step_rrf_merges_vector_and_keyword_by_chunk_id(): asyncio.run(run()) +def test_draft_steps_accumulate_by_tool_context_id(): + """Drafts are stored in app metadata and isolated by injected tool_context_id.""" + + async def run(): + app_context = ApplicationContext() + add = AddDraftStep(app_context=app_context) + read = ReadAllDraftStep(app_context=app_context) + + await add(RuntimeContext(text="first", tool_context_id="ctx-1")) + await add(RuntimeContext(text="second", tool_context_id="ctx-1")) + await add(RuntimeContext(text="other", tool_context_id="ctx-2")) + + resp = await read(RuntimeContext(tool_context_id="ctx-1")) + + assert resp.answer == "first\nsecond" + assert resp.metadata["draft_count"] == 2 + + asyncio.run(run()) + + def test_search_step_keyword_only_uses_keyword_scores_and_min_score(): """When vector has no hits, SearchStep returns keyword results directly and applies min_score.""" @@ -178,7 +199,7 @@ def test_search_step_tool_context_seen_chunks_expire_after_ttl(): step = SearchStep( file_store=store, expand_links=False, - tool_context_chunk_ttl_hours=1, + seen_ttl_hours=1, clock=lambda: now, )