refactor(Agentscope2.0): llm & embedding & agent (#271)
Some checks are pending
Pre-commit / run (ubuntu-latest) (push) Waiting to run
Tests ReMe / Unit Tests - py3.11 (push) Waiting to run
Tests ReMe / Unit Tests - py3.12 (push) Waiting to run
Tests ReMe / Unit Tests - py3.13 (push) Waiting to run

* refactor(embedding): replace embedding model with embedding store architecture

- Remove as_token_counter component and its estimated token counter implementation
- Replace BaseEmbeddingModel with BaseEmbedding that wraps AgentScope embedding models
- Add support for multiple embedding providers (OpenAI, DashScope, Gemini, Ollama)
- Introduce BaseEmbeddingStore and LocalEmbeddingStore for caching and persistence
- Update component registry to use new embedding and embedding_store types
- Modify file stores to use embedding_store instead of embedding_model
- Update health check to monitor embedding_store instead of embedding_model
- Change default config to use embedding_store with local backend
- Add estimate_token_count utility function to utils module

* refactor(llm): replace as_llm components with unified llm implementation

- Remove deprecated as_llm and as_llm_formatter modules
- Add new llm module with BaseLLM and provider-specific implementations
- Update component registry to use LLM instead of AS_LLM
- Replace all as_llm/as_llm_formatter references with llm in steps
- Update configuration schema to use llm instead of as_llm
- Rename integration test file from test_as_llm to test_llm
- Add proper docstrings to embedding store dimension property
- Add pylint disable comment for embedding model call
- Remove unused FormatterBase import in base_step
- Update token_utils with function docstring

* refactor(evolve): replace ReActAgent with Agent and update message handling

- Removed FlexReActAgent class and direct ReActAgent imports
- Updated Agent instantiation to use new constructor parameters
- Changed message content to use TextBlock format instead of plain strings
- Modified timestamp access from msg.timestamp to msg.created_at
- Updated metadata access pattern for structured outputs
- Replaced Msg.from_dict with Msg.model_validate in auto_memory.py
- Updated test mocks to patch Agent instead of ReActAgent
- Changed message serialization from to_dict to model_dump in tests
- Moved component references to base class definition
- Updated demo tools to return strings instead of ToolResponse objects

* feat(step): migrate to FunctionTool and add streaming support

- Replace deprecated ToolResponse with FunctionTool in base_step.py
- Remove unused TextBlock import from base_step.py
- Update job registration to use new FunctionTool API
- Add thinking_budget parameter to llm_demo configuration
- Introduce StreamLLMDemoStep with streaming output capability
- Add structured output support to LLMDemoStep via generate_structured_output
- Implement streaming event handling for text/thinking/tool calls
- Add integration tests for embedding functionality
- Add integration tests for structured output and streaming features
- Update tool usage in demo steps to use new function naming convention

* fix(ci): correct package installation path in unittest workflow

- Updated pip install command to use proper package path "./reme4[dev,core]"
- Fixed dependency installation step in CI workflow configuration

* chore(workflow): update python versions in unittest workflow

- Remove Python 3.10 from test matrix
- Add Python 3.11 to test matrix
- Add Python 3.12 to test matrix
- Keep Python 3.13 in test matrix
- Update matrix configuration for better version coverage

* fix(health): handle missing dimensions attribute in embedding status

- Wrap dimensions access in try-except to prevent AttributeError
- Return None when dimensions attribute is not available
- Maintain backward compatibility for components without dimensions

test(component): add comprehensive tests for BaseComponent and related classes

- Add tests for Dependency class including repr and attribute access
- Add tests for bind method with various scenarios and edge cases
- Add tests for lifecycle management and async context handling
- Add tests for standalone and context-bound dependency resolution
- Add tests for ComponentMixin path utilities

test(common): update LocalFileStore initialization parameter

- Change embedding_model parameter to embedding_store in test setup
- Update all affected test files consistently

test(registry): add complete test suite for ComponentRegistry

- Add tests for register method with explicit names and defaults
- Add tests for decorator registration pattern
- Add tests for get_all method returning copies
- Add tests for unregister and clear operations
- Add tests for error handling of invalid registrations

test(job): add comprehensive tests for BaseJob and BackgroundJob

- Add tests for step resolution and exception handling
- Add tests for backoff delay calculation with jitter
- Add tests for supervisor loop restart behavior
- Add tests for task shutdown and cancellation

test(prompt): add complete test suite for PromptHandler

- Add tests for prompt loading from dictionaries and files
- Add tests for internationalization and language fallback
- Add tests for flag filtering and variable substitution
- Add tests for format validation and error handling

test(runtime): add basic tests for RuntimeContext dictionary access

- Add tests for item getting, setting and containment checks
- Add tests for missing key error handling

* feat(evolve): add permission context and agent state management

- Import PermissionContext, PermissionMode and AgentState modules
- Add state configuration with bypass permission mode to AutoDream agents
- Add state configuration with bypass permission mode to AutoMemory agents
- Implement static _to_msg method for message validation and formatting
- Refactor message processing to use the new _to_msg method
- Ensure proper content structure for text blocks in message conversion

* style(tests): update test files with linting rules and code improvements

- Add missing pylint disable directives for docstring and attribute warnings
- Replace lambda expressions with proper function definitions in test cases
- Import Path directly instead of using lambda with __import__
- Simplify assertion checks by using truthiness instead of equality to empty dict
- Remove unused imports and reorder imports consistently
- Format dictionary literals with proper indentation and line breaks
This commit is contained in:
jinliyl 2026-06-03 11:35:43 +08:00 committed by GitHub
parent 2c35d31762
commit d8086039dc
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
47 changed files with 2447 additions and 1253 deletions

View file

@ -18,7 +18,7 @@ jobs:
strategy:
fail-fast: false
matrix:
python-version: ["3.10", "3.13"]
python-version: ["3.11", "3.12", "3.13"]
steps:
- uses: actions/checkout@v4
@ -32,7 +32,7 @@ jobs:
- name: Install dependencies
run: |
python -m pip install --upgrade pip setuptools wheel
pip install -e ".[dev,core]"
pip install -e "./reme4[dev,core]"
- name: Run tests4 unit tests
run: |

View file

@ -1,10 +1,9 @@
"""Components"""
from . import as_llm
from . import as_llm_formatter
from . import as_token_counter
from . import llm
from . import client
from . import embedding
from . import embedding_store
from . import file_catalog
from . import file_graph
from . import file_parser
@ -28,11 +27,10 @@ __all__ = [
"PromptHandler",
"RuntimeContext",
# base components
"as_llm",
"as_llm_formatter",
"as_token_counter",
"llm",
"client",
"embedding",
"embedding_store",
"file_catalog",
"file_graph",
"file_parser",

View file

@ -1,53 +0,0 @@
"""AgentScope LLM model wrappers."""
from agentscope.model import AnthropicChatModel, ChatModelBase, OpenAIChatModel
from ..base_component import BaseComponent
from ..component_registry import R
from ...enumeration import ComponentEnum
class BaseAsLLM(BaseComponent):
"""Base wrapper for AgentScope chat models. Builds ``self.model`` in ``_start``."""
component_type = ComponentEnum.AS_LLM
def __init__(self, **kwargs) -> None:
super().__init__(**kwargs)
self.model: ChatModelBase | None = None
async def _close(self) -> None:
self.model = None
@R.register("openai")
class OpenAIAsLLM(BaseAsLLM):
"""OpenAI chat model wrapper."""
async def _start(self) -> None:
self.model = OpenAIChatModel(**self.kwargs)
async def _close(self) -> None:
if self.model is not None:
assert isinstance(self.model, OpenAIChatModel)
await self.model.client.close()
@R.register("anthropic")
class AnthropicAsLLM(BaseAsLLM):
"""Anthropic chat model wrapper."""
async def _start(self) -> None:
self.model = AnthropicChatModel(**self.kwargs)
async def _close(self) -> None:
if self.model is not None:
assert isinstance(self.model, AnthropicChatModel)
await self.model.client.close()
__all__ = [
"BaseAsLLM",
"OpenAIAsLLM",
"AnthropicAsLLM",
]

View file

@ -1,44 +0,0 @@
"""AgentScope LLM formatter wrappers."""
from agentscope.formatter import AnthropicChatFormatter, FormatterBase
from .reme_openai_chat_formatter import ReMeOpenAIChatFormatter
from ..base_component import BaseComponent
from ..component_registry import R
from ...enumeration import ComponentEnum
class BaseAsLLMFormatter(BaseComponent):
"""Base wrapper for AgentScope formatters. Builds ``self.formatter`` in ``_start``."""
component_type = ComponentEnum.AS_LLM_FORMATTER
def __init__(self, **kwargs) -> None:
super().__init__(**kwargs)
self.formatter: FormatterBase | None = None
async def _close(self) -> None:
self.formatter = None
@R.register("openai")
class AsOpenAIChatFormatter(BaseAsLLMFormatter):
"""OpenAI chat formatter wrapper (uses ReMe extensions)."""
async def _start(self) -> None:
self.formatter = ReMeOpenAIChatFormatter(**self.kwargs)
@R.register("anthropic")
class AsAnthropicChatFormatter(BaseAsLLMFormatter):
"""Anthropic chat formatter wrapper."""
async def _start(self) -> None:
self.formatter = AnthropicChatFormatter(**self.kwargs)
__all__ = [
"BaseAsLLMFormatter",
"AsOpenAIChatFormatter",
"AsAnthropicChatFormatter",
]

View file

@ -1,141 +0,0 @@
"""OpenAI chat formatter with ReMe extensions: image promotion and reasoning_content."""
import json
from typing import Any
from agentscope.formatter import OpenAIChatFormatter
# noinspection PyProtectedMember
from agentscope.formatter._openai_formatter import (
_format_openai_image_block,
_to_openai_audio_data,
)
from agentscope.message import Msg, TextBlock, ImageBlock, URLSource
def _format_openai_video_block(video_block: dict) -> dict[str, Any]:
"""Convert a video block to OpenAI ``video_url`` content."""
source = video_block["source"]
if source["type"] == "url":
url = source["url"]
elif source["type"] == "base64":
url = f"data:{source['media_type']};base64,{source['data']}"
else:
raise ValueError(f"Unsupported video source type: {source['type']}")
return {"type": "video_url", "video_url": {"url": url}}
class ReMeOpenAIChatFormatter(OpenAIChatFormatter):
"""OpenAIChatFormatter + tool-result image promotion + reasoning_content passthrough."""
async def _format(self, msgs: list[Msg]) -> list[dict[str, Any]]:
"""Format ``Msg`` list into OpenAI chat-completion message dicts."""
self.assert_list_of_msgs(msgs)
messages: list[dict] = []
i = 0
while i < len(msgs):
msg = msgs[i]
content_blocks = []
tool_calls = []
reasoning_content_blocks = []
for block in msg.get_content_blocks():
typ = block.get("type")
if typ == "text":
content_blocks.append({**block})
elif typ == "thinking":
reasoning_content_blocks.append({**block})
elif typ == "tool_use":
tool_calls.append(
{
"id": block.get("id"),
"type": "function",
"function": {
"name": block.get("name"),
"arguments": json.dumps(block.get("input", {}), ensure_ascii=False),
},
},
)
elif typ == "tool_result":
textual_output, multimodal_data = self.convert_tool_result_to_string(block["output"])
messages.append(
{
"role": "tool",
"tool_call_id": block.get("id"),
"content": textual_output,
"name": block.get("name"),
},
)
# OpenAI tool messages can't carry images; promote to a follow-up user message.
promoted_blocks = []
for url, multimodal_block in multimodal_data:
if multimodal_block["type"] == "image" and self.promote_tool_result_images:
promoted_blocks.extend(
[
TextBlock(type="text", text=f"\n- The image from '{url}': "),
ImageBlock(type="image", source=URLSource(type="url", url=url)),
],
)
if promoted_blocks:
promoted_blocks = [
TextBlock(
type="text",
text="<system-info>The following are the image contents from the tool "
f"result of '{block['name']}':",
),
*promoted_blocks,
TextBlock(type="text", text="</system-info>"),
]
msgs.insert(
i + 1,
Msg(name="user", content=promoted_blocks, role="user"),
)
elif typ == "image":
content_blocks.append(_format_openai_image_block(block))
elif typ == "audio":
# Skip assistant audio — not a valid input modality.
if msg.role == "assistant":
continue
content_blocks.append(
{
"type": "input_audio",
"input_audio": _to_openai_audio_data(block["source"]),
},
)
elif typ == "video":
# Skip assistant video — not a valid input modality.
if msg.role == "assistant":
continue
content_blocks.append(_format_openai_video_block(block))
msg_openai = {
"role": msg.role,
"name": msg.name,
"content": content_blocks or None,
}
if tool_calls:
msg_openai["tool_calls"] = tool_calls
# Merge thinking blocks into reasoning_content for compatible models.
if reasoning_content_blocks:
reasoning_msg = "\n".join(r.get("thinking", "") for r in reasoning_content_blocks)
if reasoning_msg:
msg_openai["reasoning_content"] = reasoning_msg
if msg_openai["content"] or msg_openai.get("tool_calls"):
messages.append(msg_openai)
i += 1
return messages

View file

@ -1,35 +0,0 @@
"""AgentScope token counter wrappers."""
from agentscope.token import TokenCounterBase
from .estimate_token_counter import EstimatedTokenCounter
from ..base_component import BaseComponent
from ..component_registry import R
from ...enumeration import ComponentEnum
class BaseAsTokenCounter(BaseComponent):
"""Base wrapper for AgentScope token counters. Builds ``self.token_counter`` in ``_start``."""
component_type = ComponentEnum.AS_TOKEN_COUNTER
def __init__(self, **kwargs) -> None:
super().__init__(**kwargs)
self.token_counter: TokenCounterBase | None = None
async def _close(self) -> None:
self.token_counter = None
@R.register("estimated")
class EstimatedAsTokenCounter(BaseAsTokenCounter):
"""Character-based estimated token counter — fast but approximate."""
async def _start(self) -> None:
self.token_counter = EstimatedTokenCounter(**self.kwargs)
__all__ = [
"BaseAsTokenCounter",
"EstimatedAsTokenCounter",
]

View file

@ -1,21 +0,0 @@
"""Character-based token-count estimator."""
from agentscope.token import TokenCounterBase
class EstimatedTokenCounter(TokenCounterBase):
"""Approximate token count as ``encoded_byte_len / divisor``.
Cheap proxy when exact counts aren't needed; use the model's real
tokenizer for accuracy.
"""
def __init__(self, estimate_divisor: float = 4, encoding: str = "utf-8"):
if estimate_divisor <= 0:
raise ValueError("estimate_divisor must be positive")
self.estimate_divisor: float = estimate_divisor
self.encoding: str = encoding
async def count(self, text: str, **_kwargs) -> int:
"""Estimated token count for ``text``."""
return int(len(text.encode(self.encoding)) / self.estimate_divisor + 0.5)

View file

@ -1,6 +1,93 @@
"""Embedding model implementations."""
"""AgentScope embedding model wrappers."""
from .base_embedding_model import BaseEmbeddingModel
from .openai_embedding_model import OpenAIEmbeddingModel
from agentscope.embedding import (
DashScopeMultiModalEmbedding as _AsDashScopeMultiModalEmbedding,
DashScopeTextEmbedding,
EmbeddingModelBase,
GeminiTextEmbedding,
OllamaTextEmbedding,
OpenAITextEmbedding,
)
__all__ = ["BaseEmbeddingModel", "OpenAIEmbeddingModel"]
from ..base_component import BaseComponent
from ..component_registry import R
from ...enumeration import ComponentEnum
class BaseEmbedding(BaseComponent):
"""Base wrapper for AgentScope embedding models. Builds ``self.model`` in ``_start``."""
component_type = ComponentEnum.EMBEDDING
def __init__(self, **kwargs) -> None:
super().__init__(**kwargs)
self.model: EmbeddingModelBase | None = None
@property
def dimensions(self) -> int:
"""Return the embedding dimension size."""
assert self.model is not None
return self.model.dimensions
async def __call__(self, text: list[str], **kwargs) -> list[list[float]]:
assert self.model is not None
response = await self.model(text, **kwargs) # pylint: disable=not-callable
return response.embeddings
async def _close(self) -> None:
self.model = None
@R.register("openai")
class OpenAIEmbedding(BaseEmbedding):
"""OpenAI embedding model wrapper."""
async def _start(self) -> None:
self.model = OpenAITextEmbedding(**self.kwargs)
async def _close(self) -> None:
if self.model is not None:
assert isinstance(self.model, OpenAITextEmbedding)
await self.model.client.close()
@R.register("dashscope")
class DashScopeEmbedding(BaseEmbedding):
"""DashScope text embedding model wrapper."""
async def _start(self) -> None:
self.model = DashScopeTextEmbedding(**self.kwargs)
@R.register("dashscope_multimodal")
class DashScopeMultiModalEmbedding(BaseEmbedding):
"""DashScope multimodal embedding model wrapper."""
async def _start(self) -> None:
self.model = _AsDashScopeMultiModalEmbedding(**self.kwargs)
@R.register("gemini")
class GeminiEmbedding(BaseEmbedding):
"""Gemini embedding model wrapper."""
async def _start(self) -> None:
self.model = GeminiTextEmbedding(**self.kwargs)
@R.register("ollama")
class OllamaEmbedding(BaseEmbedding):
"""Ollama embedding model wrapper."""
async def _start(self) -> None:
self.model = OllamaTextEmbedding(**self.kwargs)
__all__ = [
"BaseEmbedding",
"OpenAIEmbedding",
"DashScopeEmbedding",
"DashScopeMultiModalEmbedding",
"GeminiEmbedding",
"OllamaEmbedding",
]

View file

@ -1,52 +0,0 @@
"""OpenAI-compatible async embedding model."""
from openai import AsyncOpenAI
from .base_embedding_model import BaseEmbeddingModel
from ..component_registry import R
@R.register("openai")
class OpenAIEmbeddingModel(BaseEmbeddingModel):
"""Embedding model backed by any OpenAI-compatible API."""
def __init__(self, **kwargs):
super().__init__(**kwargs)
self._client: AsyncOpenAI | None = None
async def _start(self) -> None:
"""Initialize async OpenAI client."""
self._client = AsyncOpenAI(api_key=self.api_key, base_url=self.base_url, **self.kwargs)
await super()._start()
async def _close(self) -> None:
"""Close the async OpenAI client."""
if self._client:
await self._client.close()
self._client = None
await super()._close()
async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float] | None]:
"""Call the embeddings API and return results aligned to input order."""
if self._client is None:
raise RuntimeError("Client not initialized. Call _start() first.")
create_kwargs: dict = {"model": self.model_name, "input": input_text, **kwargs}
if self.pass_dimensions:
create_kwargs["dimensions"] = self.dimensions
completion = await self._client.embeddings.create(**create_kwargs)
# Map API results back to input order
result: list[list[float] | None] = [None] * len(input_text)
for emb in completion.data:
if 0 <= emb.index < len(input_text):
vec = emb.embedding or getattr(emb, "dense_embedding", None)
if vec is not None:
result[emb.index] = list(vec)
else:
self.logger.warning(f"Empty embedding at index {emb.index}")
else:
self.logger.warning(f"Index {emb.index} out of range for input length {len(input_text)}")
return result

View file

@ -0,0 +1,6 @@
"""Embedding store implementations."""
from .base_embedding_store import BaseEmbeddingStore
from .local_embedding_store import LocalEmbeddingStore
__all__ = ["BaseEmbeddingStore", "LocalEmbeddingStore"]

View file

@ -0,0 +1,54 @@
"""Base embedding store with abstract interface for caching and retrieval."""
from abc import abstractmethod
import numpy as np
from ..base_component import BaseComponent
from ...enumeration import ComponentEnum
from ...schema import EmbNode
class BaseEmbeddingStore(BaseComponent):
"""Abstract embedding store interface.
Subclasses implement caching, persistence, and delegate actual embedding
computation to a bound ``embedding`` component.
"""
component_type = ComponentEnum.EMBEDDING_STORE
def __init__(
self,
max_batch_size: int = 10,
max_input_length: int = 8192,
max_retries: int = 3,
**kwargs,
):
super().__init__(**kwargs)
self.max_batch_size = max_batch_size
self.max_input_length = max_input_length
self.max_retries = max_retries
self.is_healthy: bool = True
@abstractmethod
async def health_check(self, timeout: float = 2.0) -> bool:
"""Probe the provider; sets and returns is_healthy."""
async def get_embedding(self, input_text: str, **kwargs) -> np.ndarray | None:
"""Embed a single text; returns None if the provider yields nothing."""
results = await self.get_embeddings([input_text], **kwargs)
return results[0] if results else None
@abstractmethod
async def get_embeddings(self, input_text: list[str], **kwargs) -> list[np.ndarray | None]:
"""Get embeddings for texts."""
async def get_node_embeddings(self, nodes: list[EmbNode], **kwargs) -> list[EmbNode]:
"""Embed each node's text in-place and return the same list."""
embeddings = await self.get_embeddings([n.text for n in nodes], **kwargs)
if len(embeddings) == len(nodes):
for node, vec in zip(nodes, embeddings):
if vec is not None:
node.embedding = vec
return nodes

View file

@ -1,73 +1,64 @@
"""Base embedding model with LRU cache and disk persistence."""
"""Local embedding store with LRU cache and disk persistence."""
import asyncio
import hashlib
import os
from abc import abstractmethod
from collections import OrderedDict
from pathlib import Path
import numpy as np
from ..base_component import BaseComponent
from ...enumeration import ComponentEnum
from ...schema import EmbNode
from .base_embedding_store import BaseEmbeddingStore
from ..component_registry import R
from ..embedding import BaseEmbedding
Miss = tuple[int, str, str] # (result_index, text, cache_key)
class BaseEmbeddingModel(BaseComponent):
"""Embedding model with LRU cache, disk persistence, and serial batching."""
@R.register("local")
class LocalEmbeddingStore(BaseEmbeddingStore):
"""Embedding store with LRU cache, disk persistence, and serial batching.
component_type = ComponentEnum.EMBEDDING_MODEL
Delegates actual embedding computation to a bound ``embedding`` component.
"""
def __init__(
self,
api_key: str | None = None,
base_url: str | None = None,
model_name: str = "",
dimensions: int = 1024,
pass_dimensions: bool = False,
max_batch_size: int = 10,
max_input_length: int = 8192,
embedding: str = "default",
max_cache_size: int = 10000,
enable_cache: bool = True,
cache_version: str = "v1",
max_retries: int = 3,
**kwargs,
):
super().__init__(**kwargs)
self.api_key = api_key or os.environ.get("EMBEDDING_API_KEY", "")
self.base_url = base_url or os.environ.get("EMBEDDING_BASE_URL", "")
self.model_name = model_name
self.dimensions = dimensions
self.pass_dimensions = pass_dimensions
self.max_batch_size = max_batch_size
self.max_input_length = max_input_length
self.embedding = self.bind(embedding, BaseEmbedding, optional=False)
self.max_cache_size = max_cache_size
self.enable_cache = enable_cache
self.cache_version = cache_version
self.max_retries = max_retries
self._cache: OrderedDict[str, np.ndarray] = OrderedDict()
self._key_suffix = f"|{model_name}|{dimensions}".encode()
self.is_healthy: bool = True
self._key_suffix: bytes = b""
@property
def dimensions(self) -> int:
"""Return the embedding dimension size."""
assert self.embedding is not None, "embedding component not bound"
return self.embedding.dimensions
@property
def cache_path(self) -> Path:
"""Path of the persisted embedding cache, namespaced by name and version."""
return self.vault_metadata_path / "embedding_cache" / f"{self.name}_{self.cache_version}.npz"
"""Return the path to the disk cache file."""
return self.component_metadata_path / f"{self.name}_{self.cache_version}.npz"
async def _start(self) -> None:
self._key_suffix = f"|{self.dimensions}".encode()
await self.load()
async def _close(self) -> None:
await self.dump()
async def health_check(self, timeout: float = 2.0) -> bool:
"""Probe the provider; sets and returns is_healthy."""
tag = f"[EMBEDDING HEALTH CHECK] name={self.name} model={self.model_name}"
tag = f"[EMBEDDING HEALTH CHECK] name={self.name}"
try:
result = await asyncio.wait_for(self._get_embeddings(["ping"]), timeout=timeout)
result = await asyncio.wait_for(self.embedding(["ping"]), timeout=timeout)
if not result or result[0] is None:
raise RuntimeError("empty embedding")
self.is_healthy = True
@ -82,39 +73,19 @@ class BaseEmbeddingModel(BaseComponent):
# -- Public API --
async def get_embedding(self, input_text: str, **kwargs) -> np.ndarray | None:
"""Embed a single text; returns None if the provider yields nothing."""
results = await self.get_embeddings([input_text], **kwargs)
return results[0] if results else None
async def get_embeddings(self, input_text: list[str], **kwargs) -> list[np.ndarray | None]:
"""Get embeddings for texts. Cache hits return immediately; misses run in serial batches."""
texts = [self._truncate(t) for t in input_text]
results, misses = self._partition_by_cache(texts)
if misses:
await self._fill_misses(misses, results, **kwargs)
return results
async def get_node_embeddings(self, nodes: list[EmbNode], **kwargs) -> list[EmbNode]:
"""Embed each node's text in-place and return the same list."""
embeddings = await self.get_embeddings([n.text for n in nodes], **kwargs)
if len(embeddings) == len(nodes):
for node, vec in zip(nodes, embeddings):
if vec is not None:
node.embedding = vec
return nodes
@abstractmethod
async def _get_embeddings(self, input_text: list[str], **kwargs) -> list[list[float] | None]:
"""Get raw embeddings from the underlying provider."""
# -- Batching --
def _truncate(self, text: str) -> str:
return text if len(text) <= self.max_input_length else text[: self.max_input_length]
def _partition_by_cache(self, texts: list[str]) -> tuple[list[np.ndarray | None], list[Miss]]:
"""Split texts into pre-filled results (hits) and a miss list to compute."""
results: list[np.ndarray | None] = [None] * len(texts)
misses: list[Miss] = []
for idx, text in enumerate(texts):
@ -127,7 +98,6 @@ class BaseEmbeddingModel(BaseComponent):
return results, misses
async def _fill_misses(self, misses: list[Miss], results: list[np.ndarray | None], **kwargs) -> None:
"""Compute miss embeddings in serial batches and write into results + cache."""
size = self.max_batch_size
batches = [misses[i : i + size] for i in range(0, len(misses), size)]
for batch in batches:
@ -136,7 +106,6 @@ class BaseEmbeddingModel(BaseComponent):
self._cache_put(key, emb)
async def _compute_batch(self, batch: list[Miss], **kwargs) -> list[tuple[int, str, np.ndarray]]:
"""Call provider for one batch with retry; returns [(idx, key, embedding)]."""
texts = [text for _, text, _ in batch]
embeddings = await self._call_with_retry(texts, **kwargs)
if not embeddings or len(embeddings) != len(texts):
@ -150,10 +119,9 @@ class BaseEmbeddingModel(BaseComponent):
return out
async def _call_with_retry(self, texts: list[str], **kwargs) -> list[list[float] | None] | None:
"""Call provider with exponential backoff on transient errors."""
for attempt in range(self.max_retries):
try:
result = await self._get_embeddings(texts, **kwargs)
result = await self.embedding(texts, **kwargs)
if result and len(result) == len(texts):
return result
except (TimeoutError, ConnectionError, OSError):
@ -197,7 +165,6 @@ class BaseEmbeddingModel(BaseComponent):
# -- Persistence --
async def load(self) -> None:
"""Load cached embeddings from disk (npz); replaces in-memory cache."""
self._cache.clear()
if not self.enable_cache or not self.cache_path.exists():
return
@ -219,7 +186,6 @@ class BaseEmbeddingModel(BaseComponent):
self.logger.info(f"Loaded {len(self._cache)} embeddings from {self.cache_path}")
async def dump(self) -> None:
"""Persist in-memory cache to disk (npz)."""
if not self.enable_cache or not self._cache:
return
await asyncio.to_thread(self._dump_sync)

View file

@ -55,7 +55,7 @@ class FaissLocalFileStore(LocalFileStore):
@property
def _dim(self) -> int:
return self.embedding_model.dimensions if self.embedding_model is not None else 0
return self.embedding_store.dimensions if self.embedding_store is not None else 0
def _new_index(self):
return self._faiss.IndexFlatIP(self._dim)
@ -111,7 +111,7 @@ class FaissLocalFileStore(LocalFileStore):
async def load(self) -> None:
"""Load chunks via the parent, then attach FAISS state (sidecar or rebuild)."""
await super().load()
if self.embedding_model is None or self._dim == 0:
if self.embedding_store is None or self._dim == 0:
self._faiss_index = None
return
if not await self._try_load_sidecar():
@ -147,7 +147,7 @@ class FaissLocalFileStore(LocalFileStore):
async def dump(self) -> None:
"""Persist chunks JSONL via the parent, then write the FAISS sidecar atomically."""
await super().dump()
if self._faiss_index is None or self.embedding_model is None:
if self._faiss_index is None or self.embedding_store is None:
return
try:
self._compact_if_needed()
@ -180,7 +180,7 @@ class FaissLocalFileStore(LocalFileStore):
}
await super().upsert(files)
if self._faiss_index is None or self.embedding_model is None:
if self._faiss_index is None or self.embedding_store is None:
return
self._sync_index_after_upsert(files, old_ids_by_path)
@ -220,7 +220,7 @@ class FaissLocalFileStore(LocalFileStore):
async def clear(self) -> None:
await super().clear()
self._faiss_index = self._new_index() if self.embedding_model is not None else None
self._faiss_index = self._new_index() if self.embedding_store is not None else None
self._id_map = []
self._id_to_row = {}
self._tombstones.clear()
@ -230,13 +230,13 @@ class FaissLocalFileStore(LocalFileStore):
# -- search -----------------------------------------------------------
async def vector_search(self, query: str, limit: int, search_filter: dict) -> list[FileChunk]:
if self.embedding_model is None or not query or self._faiss_index is None:
if self.embedding_store is None or not query or self._faiss_index is None:
return []
if self._faiss_index.ntotal == 0:
return []
try:
query_embedding = await self.embedding_model.get_embedding(query)
query_embedding = await self.embedding_store.get_embedding(query)
except Exception as e:
self._disable_embedding(f"search: {type(e).__name__}: {e}")
return []

View file

@ -5,7 +5,7 @@ import numpy as np
from .base_file_store import BaseFileStore
from ..component_registry import R
from ..embedding import BaseEmbeddingModel
from ..embedding_store import BaseEmbeddingStore
from ..file_graph import BaseFileGraph
from ..keyword_index import BaseKeywordIndex
from ...enumeration import LinkScopeEnum
@ -17,7 +17,7 @@ from ...utils import batch_cosine_similarity
class LocalFileStore(BaseFileStore):
"""In-memory file store with deferred JSONL persistence.
Composes three subcomponents: ``embedding_model`` for vector retrieval,
Composes three subcomponents: ``embedding_store`` for vector retrieval,
``keyword_index`` for full-text retrieval, and ``file_graph`` for node / link
storage. ``file_graph`` is mandatory; at least one of embedding / keyword
must be present.
@ -25,7 +25,7 @@ class LocalFileStore(BaseFileStore):
def __init__(
self,
embedding_model: str = "default",
embedding_store: str = "default",
keyword_index: str = "default",
file_graph: str = "default",
encoding: str = "utf-8",
@ -33,16 +33,16 @@ class LocalFileStore(BaseFileStore):
**kwargs,
):
super().__init__(**kwargs)
from ..embedding import OpenAIEmbeddingModel
from ..embedding_store import LocalEmbeddingStore
from ..file_graph import LocalFileGraph
from ..keyword_index import BM25Index
if not embedding_model and not keyword_index:
raise ValueError("At least one of embedding_model or keyword_index must be set.")
if not embedding_store and not keyword_index:
raise ValueError("At least one of embedding_store or keyword_index must be set.")
if not file_graph:
raise ValueError("file_graph is required for LocalFileStore.")
self.embedding_model = self.bind(embedding_model, BaseEmbeddingModel, default_factory=OpenAIEmbeddingModel)
self.embedding_store = self.bind(embedding_store, BaseEmbeddingStore, default_factory=LocalEmbeddingStore)
self.keyword_index = self.bind(keyword_index, BaseKeywordIndex, default_factory=BM25Index)
self.file_graph = self.bind(file_graph, BaseFileGraph, default_factory=LocalFileGraph)
@ -56,9 +56,9 @@ class LocalFileStore(BaseFileStore):
async def _start(self) -> None:
self.component_metadata_path.mkdir(parents=True, exist_ok=True)
await super()._start()
if self.embedding_model is not None and not await self.embedding_model.health_check():
if self.embedding_store is not None and not await self.embedding_store.health_check():
self.logger.warning(f"{self.name}: embedding unhealthy, vector disabled")
self.embedding_model = None
self.embedding_store = None
await self.load()
async def _close(self) -> None:
@ -68,10 +68,10 @@ class LocalFileStore(BaseFileStore):
def _disable_embedding(self, reason: str) -> None:
"""Drop embedding after a runtime failure; keyword search still works."""
if self.embedding_model is None:
if self.embedding_store is None:
return
self.logger.error(f"{self.name}: embedding disabled, {reason}")
self.embedding_model = None
self.embedding_store = None
# -- persistence ----------------------------------------------------------
@ -148,7 +148,7 @@ class LocalFileStore(BaseFileStore):
a new chunk reusing the same id avoids a redundant embedding call.
"""
cached: dict[str, np.ndarray] = {}
if not (old_node and self.embedding_model):
if not (old_node and self.embedding_store):
return cached
for cid in old_node.chunk_ids:
old = self.file_chunks.pop(cid, None)
@ -162,7 +162,7 @@ class LocalFileStore(BaseFileStore):
cached: dict[str, np.ndarray],
needs_embed: list[FileChunk],
) -> None:
if not self.embedding_model or chunk.embedding is not None:
if not self.embedding_store or chunk.embedding is not None:
return
if chunk.id in cached:
chunk.embedding = cached[chunk.id]
@ -170,10 +170,10 @@ class LocalFileStore(BaseFileStore):
needs_embed.append(chunk)
async def _embed_pending(self, chunks: list[FileChunk]) -> None:
if not (chunks and self.embedding_model):
if not (chunks and self.embedding_store):
return
try:
await self.embedding_model.get_node_embeddings(chunks)
await self.embedding_store.get_node_embeddings(chunks)
except Exception as e:
self._disable_embedding(f"upsert: {type(e).__name__}: {e}")
@ -221,11 +221,11 @@ class LocalFileStore(BaseFileStore):
# -- search ---------------------------------------------------------------
async def vector_search(self, query: str, limit: int, search_filter: dict) -> list[FileChunk]:
if self.embedding_model is None or not query:
if self.embedding_store is None or not query:
return []
try:
query_embedding = await self.embedding_model.get_embedding(query)
query_embedding = await self.embedding_store.get_embedding(query)
except Exception as e:
self._disable_embedding(f"search: {type(e).__name__}: {e}")
return []

View file

@ -0,0 +1,112 @@
"""LLM model wrappers for AgentScope."""
from agentscope.credential import (
AnthropicCredential,
CredentialBase,
DashScopeCredential,
DeepSeekCredential,
GeminiCredential,
MoonshotCredential,
OllamaCredential,
OpenAICredential,
XAICredential,
)
from agentscope.model import ChatModelBase
from ..base_component import BaseComponent
from ..component_registry import R
from ...enumeration import ComponentEnum
class BaseLLM(BaseComponent):
"""Base wrapper for AgentScope chat models.
Subclasses set ``credential_cls`` and inherit ``_start`` / ``_close``.
"""
component_type = ComponentEnum.LLM
credential_cls: type[CredentialBase]
def __init__(self, **kwargs) -> None:
super().__init__(**kwargs)
self.model: ChatModelBase | None = None
async def _start(self) -> None:
kwargs = dict(self.kwargs)
credential = self.credential_cls(**kwargs.pop("credential", {}))
model_cls = credential.get_chat_model_class()
params_dict = kwargs.pop("parameters", None)
parameters = model_cls.Parameters(**params_dict) if params_dict else None
self.model = model_cls(credential=credential, parameters=parameters, **kwargs)
async def _close(self) -> None:
self.model = None
@R.register("openai")
class OpenAILLM(BaseLLM):
"""OpenAI chat model wrapper."""
credential_cls = OpenAICredential
@R.register("anthropic")
class AnthropicLLM(BaseLLM):
"""Anthropic chat model wrapper."""
credential_cls = AnthropicCredential
@R.register("dashscope")
class DashScopeLLM(BaseLLM):
"""DashScope chat model wrapper."""
credential_cls = DashScopeCredential
@R.register("deepseek")
class DeepSeekLLM(BaseLLM):
"""DeepSeek chat model wrapper."""
credential_cls = DeepSeekCredential
@R.register("gemini")
class GeminiLLM(BaseLLM):
"""Gemini chat model wrapper."""
credential_cls = GeminiCredential
@R.register("moonshot")
class MoonshotLLM(BaseLLM):
"""Moonshot chat model wrapper."""
credential_cls = MoonshotCredential
@R.register("ollama")
class OllamaLLM(BaseLLM):
"""Ollama chat model wrapper."""
credential_cls = OllamaCredential
@R.register("xai")
class XAILLM(BaseLLM):
"""xAI chat model wrapper."""
credential_cls = XAICredential
__all__ = [
"BaseLLM",
"OpenAILLM",
"AnthropicLLM",
"DashScopeLLM",
"DeepSeekLLM",
"GeminiLLM",
"MoonshotLLM",
"OllamaLLM",
"XAILLM",
]

View file

@ -1,12 +1,12 @@
service:
backend: http
vault_dir: .reme
daily_dir: daily
digest_dir: digest
resource_dir: ""
# language: zh
service:
backend: http
jobs:
update_store_index_loop:
backend: background
@ -436,7 +436,7 @@ components:
default:
backend: regex
embedding_model:
embedding:
default:
backend: ${EMBEDDING_BACKEND:-openai}
api_key: ${EMBEDDING_API_KEY:-}
@ -444,17 +444,26 @@ components:
model_name: ${EMBEDDING_MODEL_NAME:-text-embedding-v4}
dimensions: 1024
as_llm:
embedding_store:
default:
backend: ${LLM_BACKEND:-anthropic}
model_name: ${LLM_MODEL_NAME:-glm-5}
api_key: ${LLM_API_KEY:-}
client_kwargs:
base_url: ${LLM_BASE_URL:-https://dashscope.aliyuncs.com/apps/anthropic}
backend: local
embedding: default
as_llm_formatter:
llm:
default:
backend: ${LLM_BACKEND:-anthropic}
model: ${LLM_MODEL_NAME:-glm-5.1}
stream: true
context_size: 200000
max_retries: 3
retry_delay: 1.0
credential:
api_key: ${LLM_API_KEY:-}
base_url: ${LLM_BASE_URL:-https://dashscope.aliyuncs.com/apps/anthropic}
parameters:
max_tokens: 100000
thinking_enable: true
thinking_budget: 38000
file_graph:
default:
@ -479,7 +488,7 @@ components:
default:
backend: local
store_name: local
# embedding_model: default
embedding_model: ""
# embedding_store: default
embedding_store: ""
keyword_index: default
file_graph: default

View file

@ -8,13 +8,11 @@ class ComponentEnum(str, Enum):
BASE = "base"
AS_LLM = "as_llm"
LLM = "llm"
AS_LLM_FORMATTER = "as_llm_formatter"
EMBEDDING = "embedding"
AS_TOKEN_COUNTER = "as_token_counter"
EMBEDDING_MODEL = "embedding_model"
EMBEDDING_STORE = "embedding_store"
FILE_PARSER = "file_parser"

View file

@ -5,14 +5,10 @@ from abc import abstractmethod, ABC
from pathlib import Path
from typing import TypeVar, TYPE_CHECKING
from agentscope.formatter import FormatterBase
from agentscope.message import TextBlock
from agentscope.model import ChatModelBase
from agentscope.token import TokenCounterBase
from agentscope.tool import Toolkit, ToolResponse
from agentscope.tool import Toolkit, FunctionTool
from ..components.base_component import ComponentMixin
from ..components.embedding import BaseEmbeddingModel
from ..components.file_parser import BaseFileParser
from ..components.file_store import BaseFileStore
from ..components.prompt_handler import PromptHandler
@ -35,7 +31,7 @@ class Ref:
Replaces the ``@property`` + ``_resolve()`` boilerplate with a single
class-level declaration::
as_llm = Ref(ChatModelBase, ComponentEnum.AS_LLM, "model")
llm = Ref(ChatModelBase, ComponentEnum.LLM, "model")
file_store = Ref(BaseFileStore, ComponentEnum.FILE_STORE)
Resolution follows a 3-source fallback identical to the old ``_resolve``:
@ -105,6 +101,9 @@ class BaseStep(ComponentMixin, ABC):
component_type = ComponentEnum.STEP
llm: ChatModelBase = Ref(ChatModelBase, ComponentEnum.LLM, "model")
file_store: BaseFileStore = Ref(BaseFileStore, ComponentEnum.FILE_STORE)
def __new__(cls, *args, **kwargs):
# Snapshot init args so copy() can rebuild an equivalent instance later.
instance = object.__new__(cls)
@ -140,14 +139,6 @@ class BaseStep(ComponentMixin, ABC):
self.prompt.load_prompt_by_class(cls)
self.prompt.load_prompt_dict(prompt_dict)
# ----- Component references (resolved lazily on first access) ----------
as_llm: ChatModelBase = Ref(ChatModelBase, ComponentEnum.AS_LLM, "model")
as_llm_formatter: FormatterBase = Ref(FormatterBase, ComponentEnum.AS_LLM_FORMATTER, "formatter")
as_token_counter: TokenCounterBase = Ref(TokenCounterBase, ComponentEnum.AS_TOKEN_COUNTER, "token_counter")
file_store: BaseFileStore = Ref(BaseFileStore, ComponentEnum.FILE_STORE)
embedding: BaseEmbeddingModel = Ref(BaseEmbeddingModel, ComponentEnum.EMBEDDING_MODEL)
@abstractmethod
async def execute(self):
"""Run the step's logic against ``self.context``."""
@ -226,20 +217,15 @@ class BaseStep(ComponentMixin, ABC):
if job is None:
raise RuntimeError(f"Job {job_name} not found")
async def run_job(**_kwargs) -> ToolResponse:
async def run_job(**_kwargs) -> str:
response = await job(**{**_kwargs, **kwargs})
return ToolResponse(content=[TextBlock(type="text", text=response.answer)])
return response.answer
toolkit.register_tool_function(
tool_func=run_job,
func_name=job_name,
func_description=job.description,
json_schema={
"type": "function",
"function": {
"name": job_name,
"description": job.description,
"parameters": job.parameters,
},
},
tool = FunctionTool(
func=run_job,
name=job_name,
description=job.description,
)
if job.parameters:
tool.input_schema = job.parameters
toolkit.tool_groups[0].tools.append(tool)

View file

@ -57,11 +57,15 @@ def _mb_str(*objs) -> str:
def _embedding_status(comp) -> dict:
cache = getattr(comp, "_embedding_cache", {}) or {}
try:
dims = comp.dimensions
except Exception:
dims = None
return {
"is_started": comp.is_started,
"is_healthy": getattr(comp, "is_healthy", None),
"model_name": getattr(comp, "model_name", None),
"dimensions": getattr(comp, "dimensions", None),
"dimensions": dims,
"cache_size": len(cache),
"memory": _mb_str(cache),
}
@ -124,7 +128,7 @@ def _keyword_index_status(comp) -> dict:
_HANDLERS = {
ComponentEnum.EMBEDDING_MODEL: _embedding_status,
ComponentEnum.EMBEDDING_STORE: _embedding_status,
ComponentEnum.FILE_GRAPH: _file_graph_status,
ComponentEnum.FILE_STORE: _file_store_status,
ComponentEnum.KEYWORD_INDEX: _keyword_index_status,
@ -140,7 +144,7 @@ def _is_healthy(ctype: ComponentEnum, status: dict) -> bool:
"""Unstarted = unhealthy; embedding model also requires is_healthy != False."""
if not status.get("is_started"):
return False
if ctype is ComponentEnum.EMBEDDING_MODEL and status.get("is_healthy") is False:
if ctype is ComponentEnum.EMBEDDING_STORE and status.get("is_healthy") is False:
return False
return True

View file

@ -1,34 +1,38 @@
"""Demo step that drives a ReActAgent via BaseStep.as_llm/as_llm_formatter."""
"""Demo step that drives an Agent via BaseStep.llm."""
from agentscope.agent import ReActAgent
from typing import Type
from agentscope.agent import Agent
from agentscope.state import AgentState
from agentscope.message import Msg, TextBlock
from agentscope.tool import Toolkit, ToolResponse
from agentscope.permission import PermissionContext, PermissionMode
from agentscope.tool import FunctionTool, Toolkit
from pydantic import BaseModel
from ..base_step import BaseStep
from ...components import R
def _add(a: float, b: float) -> ToolResponse:
def add(a: float, b: float) -> str:
"""Add two numbers and return the sum.
Args:
a: first addend
b: second addend
"""
return ToolResponse(content=[TextBlock(type="text", text=str(a + b))])
return str(a + b)
@R.register("llm_demo_step")
class LLMDemoStep(BaseStep):
"""Drive a ReActAgent powered by ``self.as_llm`` / ``self.as_llm_formatter``.
"""Drive an Agent powered by ``self.llm``.
Inputs (from RuntimeContext):
query (str, required): user message content.
sys_prompt (str, optional): system prompt for the agent.
use_add_tool (bool, optional): register the ``add`` tool when True.
console_enabled (bool, optional): mirror agent output to stdout.
Output (written to context.response.answer):fa
Output (written to context.response.answer):
The agent's final reply text.
"""
@ -39,32 +43,41 @@ class LLMDemoStep(BaseStep):
query: str = self.context.get("query", "")
sys_prompt: str = self.context.get("sys_prompt") or self.DEFAULT_SYS_PROMPT
use_add_tool: bool = bool(self.context.get("use_add_tool", False))
console_enabled: bool = bool(self.context.get("console_enabled", False))
structured_model: Type[BaseModel] | None = self.context.get("structured_model")
if not query:
self.context.response.success = False
self.context.response.answer = "Skipped: empty query"
return self.context.response
toolkit = Toolkit()
if use_add_tool:
toolkit.register_tool_function(_add)
toolkit = Toolkit(tools=[FunctionTool(add)]) if use_add_tool else Toolkit()
agent = ReActAgent(
agent = Agent(
name=self.name,
sys_prompt=sys_prompt,
model=self.as_llm,
formatter=self.as_llm_formatter,
system_prompt=sys_prompt,
model=self.llm,
toolkit=toolkit,
state=AgentState(
permission_context=PermissionContext(
mode=PermissionMode.BYPASS,
),
),
)
agent.set_console_output_enabled(console_enabled)
response: Msg = await agent.reply(
Msg(name="user", role="user", content=query),
Msg(name="user", role="user", content=[TextBlock(text=query)]),
)
text = (response.get_text_content() or "").strip()
self.logger.info(f"[{self.name}] response: {text!r}")
structured_content: dict | None = None
if structured_model is not None:
structured_resp = await self.llm.generate_structured_output(
agent.state.context,
structured_model=structured_model,
)
structured_content = structured_resp.content
self.context.response.success = True
self.context.response.answer = text
self.context.response.metadata.update(
@ -73,6 +86,7 @@ class LLMDemoStep(BaseStep):
"sys_prompt": sys_prompt,
"use_add_tool": use_add_tool,
"response": text,
"structured_output": structured_content,
},
)
return self.context.response

View file

@ -0,0 +1,136 @@
"""Demo step that drives an Agent via BaseStep.llm with streaming output."""
import json
from agentscope.agent import Agent
from agentscope.event import (
TextBlockDeltaEvent,
ThinkingBlockDeltaEvent,
ToolCallStartEvent,
ToolCallDeltaEvent,
ToolResultTextDeltaEvent,
ModelCallEndEvent,
ReplyStartEvent,
)
from agentscope.message import Msg, TextBlock
from agentscope.permission import PermissionContext, PermissionMode
from agentscope.state import AgentState
from agentscope.tool import FunctionTool, Toolkit
from ..base_step import BaseStep
from ...components import R
from ...enumeration import ChunkEnum
def add(a: float, b: float) -> str:
"""Add two numbers and return the sum.
Args:
a: first addend
b: second addend
"""
return str(a + b)
@R.register("stream_llm_demo_step")
class StreamLLMDemoStep(BaseStep):
"""Drive an Agent powered by ``self.llm`` with streaming output.
When streaming is enabled on the context, text/thinking/tool events are
pushed chunk-by-chunk via ``self.context.add_stream_string``.
When streaming is not enabled, falls back to non-streaming ``agent.reply``.
Inputs (from RuntimeContext):
query (str, required): user message content.
sys_prompt (str, optional): system prompt for the agent.
use_add_tool (bool, optional): register the ``add`` tool when True.
Output (written to context.response.answer):
The agent's final reply text.
"""
DEFAULT_SYS_PROMPT = "You are a concise assistant. Reply in one short sentence."
async def execute(self):
assert self.context is not None
query: str = self.context.get("query", "")
sys_prompt: str = self.context.get("sys_prompt") or self.DEFAULT_SYS_PROMPT
use_add_tool: bool = bool(self.context.get("use_add_tool", False))
if not query:
self.context.response.success = False
self.context.response.answer = "Skipped: empty query"
return self.context.response
toolkit = Toolkit(tools=[FunctionTool(add)]) if use_add_tool else Toolkit()
agent = Agent(
name=self.name,
system_prompt=sys_prompt,
model=self.llm,
toolkit=toolkit,
state=AgentState(
permission_context=PermissionContext(
mode=PermissionMode.BYPASS,
),
),
)
input_msg = Msg(name="user", role="user", content=[TextBlock(text=query)])
if self.context.stream:
text = await self._stream_reply(agent, input_msg)
else:
response: Msg = await agent.reply(input_msg)
text = (response.get_text_content() or "").strip()
self.logger.info(f"[{self.name}] response: {text!r}")
self.context.response.success = True
self.context.response.answer = text
self.context.response.metadata.update(
{
"query": query,
"sys_prompt": sys_prompt,
"use_add_tool": use_add_tool,
"response": text,
},
)
return self.context.response
async def _stream_reply(self, agent: Agent, input_msg: Msg) -> str:
"""Stream agent reply events to the context stream queue."""
assert self.context is not None
reply_msg: Msg | None = None
async for event in agent.reply_stream(input_msg):
if isinstance(event, ReplyStartEvent):
reply_msg = Msg(
id=event.reply_id,
name=event.name,
role=event.role,
content=[],
)
elif isinstance(event, TextBlockDeltaEvent):
await self.context.add_stream_string(event.delta, ChunkEnum.CONTENT)
elif isinstance(event, ThinkingBlockDeltaEvent):
await self.context.add_stream_string(event.delta, ChunkEnum.THINK)
elif isinstance(event, ToolCallStartEvent):
payload = json.dumps({"name": event.tool_call_name, "id": event.tool_call_id})
await self.context.add_stream_string(payload, ChunkEnum.TOOL_CALL)
elif isinstance(event, ToolCallDeltaEvent):
await self.context.add_stream_string(event.delta, ChunkEnum.TOOL_CALL)
elif isinstance(event, ToolResultTextDeltaEvent):
await self.context.add_stream_string(event.delta, ChunkEnum.TOOL_RESULT)
elif isinstance(event, ModelCallEndEvent):
usage = json.dumps(
{"input_tokens": event.input_tokens, "output_tokens": event.output_tokens},
)
await self.context.add_stream_string(usage, ChunkEnum.USAGE)
if reply_msg is not None:
reply_msg.append_event(event)
if reply_msg is not None:
return (reply_msg.get_text_content() or "").strip()
return ""

View file

@ -2,9 +2,7 @@
import datetime
import zoneinfo
from typing import Literal
from agentscope.agent import ReActAgent
from agentscope.message import Msg
@ -26,18 +24,6 @@ def format_history(messages: list[Msg], include_timestamp: bool = True) -> str:
if not text:
continue
speaker = msg.name or msg.role or "?"
header = f"[{speaker} @ {msg.timestamp}]" if include_timestamp else f"[{speaker}]"
header = f"[{speaker} @ {msg.created_at}]" if include_timestamp else f"[{speaker}]"
lines.append(f"{header}\n{text}")
return "\n\n".join(lines) or "(empty)"
class FlexReActAgent(ReActAgent):
"""ReActAgent subclass that allows structured output without forcing tool_choice='required'."""
async def _reasoning(
self,
tool_choice: Literal["auto", "none", "required"] | None = None,
) -> Msg:
if tool_choice == "required":
tool_choice = None
return await super()._reasoning(tool_choice)

View file

@ -52,11 +52,13 @@ import zoneinfo
from pathlib import Path
from typing import Literal
from agentscope.message import Msg
from agentscope.agent import Agent
from agentscope.message import Msg, TextBlock
from agentscope.permission import PermissionContext, PermissionMode
from agentscope.state import AgentState
from agentscope.tool import Toolkit
from pydantic import BaseModel, Field
from ._evolve import FlexReActAgent
from ..base_step import BaseStep
from ...components import R
@ -232,13 +234,11 @@ class Dreamer(BaseStep):
def __init__(
self,
toolkit: Toolkit | None = None,
console_enabled: bool = False,
timezone: str | None = None,
**kwargs,
):
super().__init__(**kwargs)
self.toolkit = toolkit
self.console_enabled = console_enabled
self.timezone = timezone
def _now(self) -> datetime.datetime:
@ -255,7 +255,7 @@ class Dreamer(BaseStep):
def _llm_available(self) -> bool:
try:
return self.as_llm is not None
return self.llm is not None
except Exception:
return False
@ -287,18 +287,21 @@ class Dreamer(BaseStep):
alongside its structured emission.
"""
toolkit = self._build_extract_toolkit()
agent = FlexReActAgent(
agent = Agent(
name="reme_dreamer_extract",
model=self.as_llm,
sys_prompt=self.prompt_format(
model=self.llm,
system_prompt=self.prompt_format(
"extract_system_prompt",
vault_dir=str(vault_dir),
buckets=", ".join(BUCKETS),
),
formatter=self.as_llm_formatter,
toolkit=toolkit,
state=AgentState(
permission_context=PermissionContext(
mode=PermissionMode.BYPASS,
),
),
)
agent.set_console_output_enabled(self.console_enabled)
user_message = self.prompt_format(
"extract_user_message",
today=self._now().strftime("%Y-%m-%d"),
@ -306,13 +309,14 @@ class Dreamer(BaseStep):
material_blob=material_blob,
)
msg = await agent.reply(
Msg(name="reme", role="user", content=user_message),
structured_model=ExtractedUnits,
Msg(name="reme", role="user", content=[TextBlock(text=user_message)]),
)
# Structured output lands in msg.metadata as a dict matching ExtractedUnits.
# Empty / missing → no sub-units (Phase 2 will skip).
meta = msg.metadata if isinstance(msg.metadata, dict) else {}
structured_resp = await self.llm.generate_structured_output(
agent.state.context,
structured_model=ExtractedUnits,
)
meta = structured_resp.content if isinstance(structured_resp.content, dict) else {}
cleaned: list[dict] = []
for raw in meta.get("units") or []:
if not isinstance(raw, dict):
@ -342,19 +346,22 @@ class Dreamer(BaseStep):
bucket = unit.get("bucket") or "wiki"
toolkit = self._build_integrate_toolkit()
digest_dir = getattr(self.app_context.app_config, "digest_dir", "")
agent = FlexReActAgent(
agent = Agent(
name=f"reme_dreamer_integrate_{unit.get('name', 'unit')}",
model=self.as_llm,
sys_prompt=self.prompt_format(
model=self.llm,
system_prompt=self.prompt_format(
f"integrate_system_prompt_{bucket}",
vault_dir=str(vault_dir),
digest_dir=digest_dir,
bucket=bucket,
),
formatter=self.as_llm_formatter,
toolkit=toolkit,
state=AgentState(
permission_context=PermissionContext(
mode=PermissionMode.BYPASS,
),
),
)
agent.set_console_output_enabled(self.console_enabled)
user_message = self.prompt_format(
"integrate_user_message",
hint=hint or "(none)",
@ -363,12 +370,14 @@ class Dreamer(BaseStep):
unit_summary=unit.get("summary", ""),
material_blob=material_blob,
)
msg = await agent.reply(
Msg(name="reme", role="user", content=user_message),
await agent.reply(
Msg(name="reme", role="user", content=[TextBlock(text=user_message)]),
)
structured_resp = await self.llm.generate_structured_output(
agent.state.context,
structured_model=IntegrateOutcome,
)
meta = msg.metadata if isinstance(msg.metadata, dict) else {}
return IntegrateOutcome.model_validate(meta)
return IntegrateOutcome.model_validate(structured_resp.content)
async def dream_one(self, path: str, hint: str = "") -> DreamResult:
"""Run the full extract + integrate pipeline on one vault-relative
@ -389,7 +398,7 @@ class Dreamer(BaseStep):
used_llm=False,
skipped=True,
path=path,
error="no as_llm configured; dreaming requires an LLM",
error="no llm configured; dreaming requires an LLM",
)
material_blob = _pack_material(self.file_store, path)

View file

@ -17,8 +17,10 @@ Output (written to context.response):
metadata: {path, created}.
"""
from agentscope.agent import ReActAgent
from agentscope.message import Msg
from agentscope.agent import Agent
from agentscope.message import Msg, TextBlock
from agentscope.permission import PermissionContext, PermissionMode
from agentscope.state import AgentState
from agentscope.tool import Toolkit
from ._evolve import format_history, now
@ -28,18 +30,23 @@ from ...components import R
@R.register("auto_memory_step")
class AutoMemoryStep(BaseStep):
"""Record conversation facts into a daily note via a ReAct agent."""
"""Record conversation facts into a daily note via an Agent."""
def __init__(self, console_enabled: bool = False, **kwargs):
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.console_enabled = console_enabled
self.agent_tools: list[str] = ["read", "edit", "frontmatter_update", "write"]
@staticmethod
def _to_msg(item) -> Msg:
if isinstance(item, Msg):
return item
if isinstance(item, dict) and isinstance(item.get("content"), str):
item = {**item, "content": [{"type": "text", "text": item["content"]}]}
return Msg.model_validate(item)
async def execute(self):
assert self.context is not None
messages: list[Msg] = [
item if isinstance(item, Msg) else Msg.from_dict(item) for item in self.context.get("messages", [])
]
messages: list[Msg] = [self._to_msg(item) for item in self.context.get("messages", [])]
session_id: str = self.context.get("session_id", "")
memory_hint: str = self.context.get("memory_hint", "")
current = now(self.context.get("timezone"))
@ -62,14 +69,17 @@ class AutoMemoryStep(BaseStep):
for job_name in self.agent_tools:
self.add_as_tool(toolkit, job_name)
agent = ReActAgent(
agent = Agent(
name="auto_memory",
model=self.as_llm,
sys_prompt=self.prompt_format("system_prompt"),
formatter=self.as_llm_formatter,
model=self.llm,
system_prompt=self.prompt_format("system_prompt"),
toolkit=toolkit,
state=AgentState(
permission_context=PermissionContext(
mode=PermissionMode.BYPASS,
),
),
)
agent.set_console_output_enabled(self.console_enabled)
template_key = "user_message_create" if created else "user_message_update"
user_message: str = self.prompt_format(
@ -81,7 +91,7 @@ class AutoMemoryStep(BaseStep):
history=format_history(messages),
)
final_msg: Msg = await agent.reply(Msg(name="reme", role="user", content=user_message))
final_msg: Msg = await agent.reply(Msg(name="reme", role="user", content=[TextBlock(text=user_message)]))
self.context.response.success = True
self.context.response.answer = (final_msg.get_text_content() or "").strip()

View file

@ -13,6 +13,7 @@ from .logger_utils import get_logger
from .logo_utils import print_logo
from .service_utils import find_reme, locate_reme, precheck_start, cli_find_reme
from .similarity_utils import cosine_similarity, batch_cosine_similarity
from .token_utils import estimate_token_count
__all__ = [
"hash_text",
@ -31,4 +32,5 @@ __all__ = [
"cli_find_reme",
"cosine_similarity",
"batch_cosine_similarity",
"estimate_token_count",
]

View file

@ -0,0 +1,10 @@
"""Token count estimation."""
def estimate_token_count(
text: str,
estimate_divisor: float = 4,
encoding: str = "utf-8",
) -> int:
"""Estimate the number of tokens in *text* by byte length."""
return int(len(text.encode(encoding)) / estimate_divisor + 0.5)

View file

@ -1,92 +0,0 @@
"""Integration tests: drive ReActAgent through LLMDemoStep + Application wiring.
Requires LLM_API_KEY (and optionally LLM_BASE_URL / LLM_MODEL_NAME) in the
environment or a .env file at the repo root. Hits the real Anthropic API.
"""
import asyncio
import os
import tempfile
from reme4 import Application
from reme4.config import resolve_app_config
from reme4.steps.common.llm_demo import LLMDemoStep
from reme4.utils import load_env
load_env()
class _temp_chdir:
"""chdir to path for the duration of the block; restore on exit."""
def __init__(self, path):
self.path = path
self._old = None
def __enter__(self):
self._old = os.getcwd()
os.chdir(self.path)
return self
def __exit__(self, *exc):
os.chdir(self._old)
async def _make_app() -> Application:
"""Build and start an Application from the default config (LLM wired via env vars)."""
cfg = resolve_app_config(log_to_console=False, log_to_file=False, enable_logo=False)
app = Application(**cfg)
await app.start()
return app
def test_llm_demo_step_basic_chat():
"""LLMDemoStep drives ReActAgent through self.as_llm/as_llm_formatter."""
async def run():
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
app = await _make_app()
try:
step = LLMDemoStep(app_context=app.context)
response = await step(
query="What is 1 + 1? Reply with just the number.",
)
text = (response.answer or "").strip()
print(f"\n[basic_chat] response: {text!r}")
assert text, "Empty assistant response"
assert "2" in text, f"Expected '2' in response, got: {text!r}"
print("✓ test_llm_demo_step_basic_chat passed")
finally:
await app.close()
asyncio.run(run())
def test_llm_demo_step_with_tool():
"""LLMDemoStep registers the add tool and the agent invokes it."""
async def run():
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
app = await _make_app()
try:
step = LLMDemoStep(app_context=app.context)
response = await step(
query="Use the add tool to compute 21 + 21 and report the result.",
sys_prompt="Use the `add` tool whenever the user asks to add numbers.",
use_add_tool=True,
)
text = (response.answer or "").strip()
print(f"\n[with_tool] response: {text!r}")
assert "42" in text, f"Expected '42' in response, got: {text!r}"
print("✓ test_llm_demo_step_with_tool passed")
finally:
await app.close()
asyncio.run(run())
if __name__ == "__main__":
print("=== LLMDemoStep + ReActAgent integration tests ===")
test_llm_demo_step_basic_chat()
test_llm_demo_step_with_tool()
print("\nAll integration tests passed!")

View file

@ -20,7 +20,7 @@ import tempfile
from datetime import date as _date
from pathlib import Path
from agentscope.agent import ReActAgent
from agentscope.agent import Agent
from reme4 import Application
from reme4.config import resolve_app_config
@ -168,7 +168,7 @@ def _read_text(p: Path) -> str:
class _AgentMemoryRecorder:
"""Monkey-patches ReActAgent.__init__ to capture every agent created inside
"""Monkey-patches Agent.__init__ to capture every agent created inside
the ``with`` block, then dumps each agent's memory to a jsonl file in
DUMP_DIR on exit.
"""
@ -177,13 +177,13 @@ class _AgentMemoryRecorder:
"""init"""
self.dump_dir = dump_dir
self.prefix = prefix
self.agents: list[ReActAgent] = []
self.agents: list[Agent] = []
self._orig_init = None
self.dumped_paths: list[Path] = []
def __enter__(self):
"""Monkey-patch ReActAgent.__init__."""
self._orig_init = ReActAgent.__init__
"""Monkey-patch Agent.__init__."""
self._orig_init = Agent.__init__
agents = self.agents
orig = self._orig_init
@ -191,25 +191,25 @@ class _AgentMemoryRecorder:
orig(agent_self, *args, **kwargs)
agents.append(agent_self)
ReActAgent.__init__ = _capturing_init
Agent.__init__ = _capturing_init
return self
def __exit__(self, *exc):
"""Restore the original __init__."""
ReActAgent.__init__ = self._orig_init
Agent.__init__ = self._orig_init
async def dump(self) -> list[Path]:
"""Dump all agent memories."""
"""Dump all agent context histories."""
for stale in self.dump_dir.glob(f"{self.prefix}_*.jsonl"):
stale.unlink()
for idx, agent in enumerate(self.agents, 1):
messages = await agent.memory.get_memory()
messages = agent.state.context
name = getattr(agent, "name", "agent") or "agent"
out_path = self.dump_dir / f"{self.prefix}_{idx:02d}_{name}.jsonl"
with out_path.open("w", encoding="utf-8") as f:
for msg in messages:
f.write(json.dumps(msg.to_dict(), ensure_ascii=False, default=str) + "\n")
f.write(json.dumps(msg.model_dump(), ensure_ascii=False, default=str) + "\n")
self.dumped_paths.append(out_path)
return self.dumped_paths

View file

@ -0,0 +1,198 @@
"""Integration tests: drive embedding store through Application wiring.
Requires EMBEDDING_API_KEY (and optionally EMBEDDING_BACKEND / EMBEDDING_BASE_URL /
EMBEDDING_MODEL_NAME) in the environment or a .env file at the repo root.
Hits the real embedding API.
"""
import asyncio
import os
import tempfile
import numpy as np
from reme4 import Application
from reme4.config import resolve_app_config
from reme4.enumeration import ComponentEnum
from reme4.schema import EmbNode
from reme4.utils import cosine_similarity, load_env
load_env()
class _temp_chdir:
"""chdir to path for the duration of the block; restore on exit."""
def __init__(self, path):
self.path = path
self._old = None
def __enter__(self):
self._old = os.getcwd()
os.chdir(self.path)
return self
def __exit__(self, *exc):
os.chdir(self._old)
async def _make_app() -> Application:
"""Build and start an Application from the default config."""
cfg = resolve_app_config(log_to_console=False, log_to_file=False, enable_logo=False)
app = Application(**cfg)
await app.start()
return app
def test_embedding_health_check():
"""health_check() returns True with a working API key."""
async def run():
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
app = await _make_app()
try:
store = app.context.components[ComponentEnum.EMBEDDING_STORE]["default"]
result = await store.health_check(timeout=10.0)
assert result is True, f"health_check returned {result}"
assert store.is_healthy is True
print("✓ test_embedding_health_check passed")
finally:
await app.close()
asyncio.run(run())
def test_embedding_single_text():
"""Single text produces a valid embedding vector."""
async def run():
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
app = await _make_app()
try:
store = app.context.components[ComponentEnum.EMBEDDING_STORE]["default"]
emb = await store.get_embedding("Hello, world!")
assert emb is not None, "get_embedding returned None"
assert emb.shape == (store.dimensions,), f"shape {emb.shape} != ({store.dimensions},)"
assert emb.dtype == np.float16, f"dtype {emb.dtype} != float16"
assert np.linalg.norm(emb) > 0, "embedding is a zero vector"
print(f"\n [single] len={len(emb)}, first5={emb[:5].tolist()}")
print("✓ test_embedding_single_text passed")
finally:
await app.close()
asyncio.run(run())
def test_embedding_multiple_texts():
"""Batch embedding returns correct count and shapes."""
async def run():
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
app = await _make_app()
try:
store = app.context.components[ComponentEnum.EMBEDDING_STORE]["default"]
texts = ["cat", "dog", "house"]
results = await store.get_embeddings(texts)
assert len(results) == 3, f"expected 3 results, got {len(results)}"
for i, emb in enumerate(results):
assert emb is not None, f"result[{i}] is None"
assert emb.shape == (store.dimensions,), f"result[{i}] shape mismatch"
assert np.linalg.norm(emb) > 0, f"result[{i}] is a zero vector"
print(f"\n [{texts[i]}] len={len(emb)}, first5={emb[:5].tolist()}")
print("✓ test_embedding_multiple_texts passed")
finally:
await app.close()
asyncio.run(run())
def test_embedding_cache_hit():
"""Same text returns cached result on second call."""
async def run():
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
app = await _make_app()
try:
store = app.context.components[ComponentEnum.EMBEDDING_STORE]["default"]
text = "test caching behavior"
emb1 = await store.get_embedding(text)
cache_size_after_first = len(store._cache) # pylint: disable=protected-access
emb2 = await store.get_embedding(text)
cache_size_after_second = len(store._cache) # pylint: disable=protected-access
assert emb1 is not None and emb2 is not None
assert cache_size_after_second == cache_size_after_first, "cache grew on second call"
assert np.array_equal(emb1, emb2), "cached embedding differs from original"
print(f"\n [cache] len={len(emb1)}, first5={emb1[:5].tolist()}")
print("✓ test_embedding_cache_hit passed")
finally:
await app.close()
asyncio.run(run())
def test_embedding_similarity():
"""Semantically similar texts have higher cosine similarity."""
async def run():
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
app = await _make_app()
try:
store = app.context.components[ComponentEnum.EMBEDDING_STORE]["default"]
text_a = "The cat sat on the mat"
text_b = "A kitten rested on the rug"
text_c = "Quantum computing uses qubits for parallel computation"
results = await store.get_embeddings([text_a, text_b, text_c])
emb_a, emb_b, emb_c = results
sim_ab = cosine_similarity(emb_a.tolist(), emb_b.tolist())
sim_ac = cosine_similarity(emb_a.tolist(), emb_c.tolist())
print(f"\n sim(cat/kitten) = {sim_ab:.4f}")
print(f" sim(cat/quantum) = {sim_ac:.4f}")
assert sim_ab > 0.4, f"similar texts sim={sim_ab:.4f}, expected > 0.4"
assert sim_ab > sim_ac, f"similar pair ({sim_ab:.4f}) not > dissimilar ({sim_ac:.4f})"
print("✓ test_embedding_similarity passed")
finally:
await app.close()
asyncio.run(run())
def test_embedding_node_embeddings():
"""get_node_embeddings fills embedding field on EmbNode objects."""
async def run():
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
app = await _make_app()
try:
store = app.context.components[ComponentEnum.EMBEDDING_STORE]["default"]
nodes = [
EmbNode(text="first node text"),
EmbNode(text="second node text"),
]
result = await store.get_node_embeddings(nodes)
assert result is nodes, "get_node_embeddings should return the same list"
for i, node in enumerate(nodes):
assert node.embedding is not None, f"node[{i}].embedding is None"
assert node.embedding.shape == (store.dimensions,), f"node[{i}] shape mismatch"
print(f"\n [node{i}] len={len(node.embedding)}, first5={node.embedding[:5].tolist()}")
print("✓ test_embedding_node_embeddings passed")
finally:
await app.close()
asyncio.run(run())
if __name__ == "__main__":
print("=== Embedding integration tests ===")
test_embedding_health_check()
test_embedding_single_text()
test_embedding_multiple_texts()
test_embedding_cache_hit()
test_embedding_similarity()
test_embedding_node_embeddings()
print("\nAll embedding integration tests passed!")

View file

@ -0,0 +1,175 @@
"""Integration tests: drive Agent through LLMDemoStep + Application wiring.
Requires LLM_API_KEY (and optionally LLM_BASE_URL / LLM_MODEL_NAME) in the
environment or a .env file at the repo root. Hits the real Anthropic API.
"""
import asyncio
import os
import tempfile
from typing import Literal
from pydantic import BaseModel, Field
from reme4 import Application
from reme4.config import resolve_app_config
from reme4.steps.common.llm_demo import LLMDemoStep
from reme4.utils import load_env
load_env()
class _temp_chdir:
"""chdir to path for the duration of the block; restore on exit."""
def __init__(self, path):
self.path = path
self._old = None
def __enter__(self):
self._old = os.getcwd()
os.chdir(self.path)
return self
def __exit__(self, *exc):
os.chdir(self._old)
async def _make_app() -> Application:
"""Build and start an Application from the default config (LLM wired via env vars)."""
cfg = resolve_app_config(log_to_console=False, log_to_file=False, enable_logo=False)
app = Application(**cfg)
await app.start()
return app
def test_llm_demo_step_basic_chat():
"""LLMDemoStep drives Agent through self.llm."""
async def run():
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
app = await _make_app()
try:
step = LLMDemoStep(app_context=app.context)
response = await step(
query="What is 1 + 1? Reply with just the number.",
)
text = (response.answer or "").strip()
print(f"\n[basic_chat] response: {text!r}")
assert text, "Empty assistant response"
assert "2" in text, f"Expected '2' in response, got: {text!r}"
print("✓ test_llm_demo_step_basic_chat passed")
finally:
await app.close()
asyncio.run(run())
def test_llm_demo_step_with_tool():
"""LLMDemoStep registers the add tool and the agent invokes it."""
async def run():
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
app = await _make_app()
try:
step = LLMDemoStep(app_context=app.context)
response = await step(
query="Use the add tool to compute 21 + 21 and report the result.",
sys_prompt="Use the `add` tool whenever the user asks to add numbers.",
use_add_tool=True,
)
text = (response.answer or "").strip()
print(f"\n[with_tool] response: {text!r}")
assert "42" in text, f"Expected '42' in response, got: {text!r}"
print("✓ test_llm_demo_step_with_tool passed")
finally:
await app.close()
asyncio.run(run())
class MathResult(BaseModel):
"""Structured output for a math computation."""
expression: str = Field(description="The math expression that was evaluated")
result: float = Field(description="The numeric result")
explanation: str = Field(description="Brief explanation of the computation")
class SentimentAnalysis(BaseModel):
"""Structured output for sentiment analysis."""
sentiment: Literal["positive", "negative", "neutral"] = Field(
description="The overall sentiment of the text",
)
confidence: float = Field(
description="Confidence score between 0 and 1",
)
key_phrases: list[str] = Field(
description="Key phrases that indicate the sentiment",
)
def test_llm_demo_step_structured_output():
"""LLMDemoStep generates structured output via generate_structured_output."""
async def run():
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
app = await _make_app()
try:
step = LLMDemoStep(app_context=app.context)
response = await step(
query="What is 15 multiplied by 7? Show your work.",
sys_prompt="You are a math tutor. Solve the problem step by step.",
structured_model=MathResult,
)
structured = response.metadata.get("structured_output")
print(f"\n[structured_output] result: {structured}")
assert structured is not None, "structured_output should not be None"
assert "result" in structured, "structured_output should have 'result' field"
assert structured["result"] == 105, f"Expected result=105, got: {structured['result']}"
assert "expression" in structured, "structured_output should have 'expression' field"
assert "explanation" in structured, "structured_output should have 'explanation' field"
print("✓ test_llm_demo_step_structured_output passed")
finally:
await app.close()
asyncio.run(run())
def test_llm_demo_step_structured_output_enum():
"""LLMDemoStep structured output with Literal/enum fields."""
async def run():
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
app = await _make_app()
try:
step = LLMDemoStep(app_context=app.context)
response = await step(
query="Analyze the sentiment: 'I absolutely love this product! It exceeded all my expectations.'",
sys_prompt="You are a sentiment analysis expert. Analyze the given text.",
structured_model=SentimentAnalysis,
)
structured = response.metadata.get("structured_output")
print(f"\n[structured_enum] result: {structured}")
assert structured is not None, "structured_output should not be None"
assert (
structured["sentiment"] == "positive"
), f"Expected sentiment='positive', got: {structured['sentiment']}"
assert 0 <= structured["confidence"] <= 1, f"Confidence should be 0-1, got: {structured['confidence']}"
assert isinstance(structured["key_phrases"], list), "key_phrases should be a list"
assert len(structured["key_phrases"]) > 0, "key_phrases should not be empty"
print("✓ test_llm_demo_step_structured_output_enum passed")
finally:
await app.close()
asyncio.run(run())
if __name__ == "__main__":
print("=== LLMDemoStep + Agent integration tests ===")
test_llm_demo_step_basic_chat()
test_llm_demo_step_with_tool()
test_llm_demo_step_structured_output()
test_llm_demo_step_structured_output_enum()
print("\nAll integration tests passed!")

View file

@ -0,0 +1,217 @@
"""Integration tests: stream Agent output through StreamLLMDemoStep.
Requires LLM_API_KEY (and optionally LLM_BASE_URL / LLM_MODEL_NAME) in the
environment or a .env file at the repo root. Hits the real LLM API.
"""
import asyncio
import os
import tempfile
from reme4 import Application
from reme4.config import resolve_app_config
from reme4.enumeration import ChunkEnum
from reme4.schema import StreamChunk
from reme4.steps.common.stream_llm_demo import StreamLLMDemoStep
from reme4.utils import load_env
load_env()
class _temp_chdir:
"""chdir to path for the duration of the block; restore on exit."""
def __init__(self, path):
self.path = path
self._old = None
def __enter__(self):
self._old = os.getcwd()
os.chdir(self.path)
return self
def __exit__(self, *exc):
os.chdir(self._old)
async def _make_app() -> Application:
"""Build and start an Application from the default config."""
cfg = resolve_app_config(log_to_console=False, log_to_file=False, enable_logo=False)
app = Application(**cfg)
await app.start()
return app
async def _test_stream_llm_basic_chat():
"""StreamLLMDemoStep streams text chunks via add_stream_string."""
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
app = await _make_app()
try:
step = StreamLLMDemoStep(app_context=app.context)
queue: asyncio.Queue = asyncio.Queue()
response = await step(
stream_queue=queue,
query="What is 1 + 1? Reply with just the number.",
)
# Collect all chunks from the queue
chunks = []
while not queue.empty():
chunks.append(await queue.get())
# Should have received CONTENT chunks
content_chunks = [c for c in chunks if c.chunk_type == ChunkEnum.CONTENT]
print(f"\n[stream_basic] got {len(content_chunks)} CONTENT chunks")
assert len(content_chunks) > 0, "Expected at least one CONTENT chunk"
# Final answer should be populated
text = (response.answer or "").strip()
print(f"[stream_basic] final answer: {text!r}")
assert text, "Empty assistant response"
assert "2" in text, f"Expected '2' in response, got: {text!r}"
# Concatenated stream text should match the final answer
streamed_text = "".join(c.chunk for c in content_chunks)
assert streamed_text.strip() == text, f"Stream text mismatch: {streamed_text!r} vs {text!r}"
print("✓ test_stream_llm_basic_chat passed")
finally:
await app.close()
async def _test_stream_llm_with_tool():
"""StreamLLMDemoStep streams tool call events when tools are used."""
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
app = await _make_app()
try:
step = StreamLLMDemoStep(app_context=app.context)
queue: asyncio.Queue = asyncio.Queue()
response = await step(
stream_queue=queue,
query="Use the add tool to compute 21 + 21 and report the result.",
sys_prompt="Use the `add` tool whenever the user asks to add numbers.",
use_add_tool=True,
)
# Collect all chunks
chunks = []
while not queue.empty():
chunks.append(await queue.get())
tool_call_chunks = [c for c in chunks if c.chunk_type == ChunkEnum.TOOL_CALL]
tool_result_chunks = [c for c in chunks if c.chunk_type == ChunkEnum.TOOL_RESULT]
content_chunks = [c for c in chunks if c.chunk_type == ChunkEnum.CONTENT]
print(f"\n[stream_tool] TOOL_CALL chunks: {len(tool_call_chunks)}")
print(f"[stream_tool] TOOL_RESULT chunks: {len(tool_result_chunks)}")
print(f"[stream_tool] CONTENT chunks: {len(content_chunks)}")
assert len(tool_call_chunks) > 0, "Expected TOOL_CALL chunks"
assert len(tool_result_chunks) > 0, "Expected TOOL_RESULT chunks"
text = (response.answer or "").strip()
print(f"[stream_tool] final answer: {text!r}")
assert "42" in text, f"Expected '42' in response, got: {text!r}"
print("✓ test_stream_llm_with_tool passed")
finally:
await app.close()
async def _test_stream_llm_fallback_no_stream():
"""Without stream_queue, falls back to non-streaming reply."""
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
app = await _make_app()
try:
step = StreamLLMDemoStep(app_context=app.context)
response = await step(
query="What is 1 + 1? Reply with just the number.",
)
text = (response.answer or "").strip()
print(f"\n[fallback] response: {text!r}")
assert text, "Empty assistant response"
assert "2" in text, f"Expected '2' in response, got: {text!r}"
print("✓ test_stream_llm_fallback_no_stream passed")
finally:
await app.close()
def test_stream_llm_basic_chat():
"""StreamLLMDemoStep streams text chunks via add_stream_string."""
asyncio.run(_test_stream_llm_basic_chat())
def test_stream_llm_with_tool():
"""StreamLLMDemoStep streams tool call events when tools are used."""
asyncio.run(_test_stream_llm_with_tool())
def test_stream_llm_fallback_no_stream():
"""Without stream_queue, falls back to non-streaming reply."""
asyncio.run(_test_stream_llm_fallback_no_stream())
async def _demo_stream_print():
"""Real-time streaming print demo — ask a longer question to see chunked output."""
import sys # pylint: disable=import-outside-toplevel,redefined-outer-name
with tempfile.TemporaryDirectory() as tmp, _temp_chdir(tmp):
app = await _make_app()
try:
step = StreamLLMDemoStep(app_context=app.context)
queue: asyncio.Queue = asyncio.Queue()
query = (
"Please explain in detail how neural networks learn through backpropagation. "
"Include the chain rule, gradient descent, and give a concrete example with numbers."
)
async def consumer():
"""Print chunks to terminal in real-time."""
while True:
chunk = await queue.get()
if chunk.done:
break
if chunk.chunk_type == ChunkEnum.CONTENT:
sys.stdout.write(chunk.chunk)
sys.stdout.flush()
elif chunk.chunk_type == ChunkEnum.THINK:
sys.stdout.write(f"\033[2m{chunk.chunk}\033[0m")
sys.stdout.flush()
elif chunk.chunk_type == ChunkEnum.TOOL_CALL:
sys.stdout.write(f"\n\033[33m[tool_call] {chunk.chunk}\033[0m")
sys.stdout.flush()
elif chunk.chunk_type == ChunkEnum.TOOL_RESULT:
sys.stdout.write(f"\033[32m{chunk.chunk}\033[0m")
sys.stdout.flush()
print()
consumer_task = asyncio.create_task(consumer())
await step(
stream_queue=queue,
query=query,
sys_prompt="You are a knowledgeable AI teacher. Explain concepts thoroughly.",
)
# Signal done so consumer exits
await queue.put(StreamChunk(chunk_type=ChunkEnum.DONE, chunk="", done=True))
await consumer_task
finally:
await app.close()
async def _run_all():
print("=== StreamLLMDemoStep integration tests ===")
await _test_stream_llm_basic_chat()
await _test_stream_llm_with_tool()
await _test_stream_llm_fallback_no_stream()
print("\nAll stream integration tests passed!")
if __name__ == "__main__":
import sys
if len(sys.argv) > 1 and sys.argv[1] == "demo":
asyncio.run(_demo_stream_print())
else:
asyncio.run(_run_all())

View file

@ -60,7 +60,7 @@ async def _make_scan_step(
suffix_filters: list[str] | None = None,
recursive: bool = True,
) -> tuple[ScanChangesStep, RuntimeContext, LocalFileStore, ChunkedFileParser]:
fs = LocalFileStore(name="test_store", embedding_model="")
fs = LocalFileStore(name="test_store", embedding_store="")
parser = ChunkedFileParser()
await fs.start()
await parser.start()

View file

@ -0,0 +1,344 @@
"""Tests for BaseComponent, Dependency, and ComponentMixin."""
# pylint: disable=protected-access,missing-function-docstring,missing-class-docstring,attribute-defined-outside-init
import asyncio
import os
import tempfile
import pytest
from reme4.components.base_component import BaseComponent, ComponentMixin, Dependency
from reme4.enumeration import ComponentEnum
# -- Test subclasses ----------------------------------------------------------
class StubComponent(BaseComponent):
component_type = ComponentEnum.FILE_PARSER
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.start_count = 0
self.close_count = 0
async def _start(self):
self.start_count += 1
async def _close(self):
self.close_count += 1
class DepTarget(BaseComponent):
component_type = ComponentEnum.KEYWORD_INDEX
class RequiredDepTarget(BaseComponent):
component_type = ComponentEnum.FILE_GRAPH
# -- Dependency ---------------------------------------------------------------
def test_dependency_repr_optional():
dep = Dependency(ComponentEnum.FILE_PARSER, "my_parser", optional=True)
assert "?" in repr(dep)
assert "file_parser" in repr(dep)
def test_dependency_repr_required():
dep = Dependency(ComponentEnum.FILE_PARSER, "my_parser", optional=False)
assert "?" not in repr(dep)
def test_dependency_getattr_raises():
dep = Dependency(ComponentEnum.FILE_PARSER, "my_parser")
with pytest.raises(RuntimeError, match="accessed before start"):
_ = dep.some_method
# -- bind ---------------------------------------------------------------------
def test_bind_returns_none_for_empty_name():
result = BaseComponent.bind(None, DepTarget)
assert result is None
result = BaseComponent.bind("", DepTarget)
assert result is None
def test_bind_returns_dependency_placeholder():
result = BaseComponent.bind("my_index", DepTarget)
assert isinstance(result, Dependency)
assert result.ctype == ComponentEnum.KEYWORD_INDEX
assert result.name == "my_index"
def test_bind_rejects_base_component_type():
class BadTarget(BaseComponent):
component_type = ComponentEnum.BASE
with pytest.raises(TypeError, match="non-BASE"):
BaseComponent.bind("x", BadTarget)
def test_bind_rejects_no_component_type():
class NoType:
pass
with pytest.raises(TypeError, match="non-BASE"):
BaseComponent.bind("x", NoType)
def test_bind_with_default_factory():
def factory():
return DepTarget(name="default")
result = BaseComponent.bind("idx", DepTarget, default_factory=factory)
assert isinstance(result, Dependency)
assert result.default_factory is factory
def test_bind_optional_flag():
dep = BaseComponent.bind("idx", DepTarget, optional=False)
assert dep.optional is False
# -- dependencies property ----------------------------------------------------
def test_dependencies_lists_unresolved():
comp = StubComponent()
comp.dep1 = Dependency(ComponentEnum.KEYWORD_INDEX, "a")
comp.dep2 = Dependency(ComponentEnum.FILE_GRAPH, "b")
comp.normal_attr = "not a dep"
deps = comp.dependencies
assert len(deps) == 2
# -- lifecycle ----------------------------------------------------------------
def test_start_close_idempotent():
async def run():
comp = StubComponent()
await comp.start()
await comp.start()
assert comp.start_count == 1
assert comp.is_started is True
await comp.close()
await comp.close()
assert comp.close_count == 1
assert comp.is_started is False
asyncio.run(run())
def test_restart():
async def run():
comp = StubComponent()
await comp.start()
await comp.restart()
assert comp.start_count == 2
assert comp.close_count == 1
assert comp.is_started is True
await comp.close()
asyncio.run(run())
def test_async_context_manager():
async def run():
comp = StubComponent()
async with comp as c:
assert c is comp
assert comp.is_started is True
assert comp.is_started is False
asyncio.run(run())
# -- standalone resolution ----------------------------------------------------
def test_resolve_standalone_optional_becomes_none():
async def run():
comp = StubComponent()
comp.dep = BaseComponent.bind("idx", DepTarget)
await comp.start()
assert comp.dep is None
await comp.close()
asyncio.run(run())
def test_resolve_standalone_with_default_factory():
async def run():
comp = StubComponent()
comp.dep = BaseComponent.bind(
"idx",
DepTarget,
default_factory=lambda: DepTarget(name="auto"),
)
await comp.start()
assert isinstance(comp.dep, DepTarget)
assert comp.dep.name == "auto"
assert comp.dep in comp._owned
await comp.close()
asyncio.run(run())
def test_resolve_standalone_required_no_factory_keeps_placeholder():
async def run():
comp = StubComponent()
comp.dep = BaseComponent.bind("idx", DepTarget, optional=False)
await comp.start()
assert isinstance(comp.dep, Dependency)
await comp.close()
asyncio.run(run())
# -- owned component lifecycle cascade ----------------------------------------
def test_owned_components_started_and_closed():
async def run():
owned = StubComponent(name="owned")
parent = StubComponent(name="parent")
parent.dep = BaseComponent.bind(
"sub",
StubComponent,
default_factory=lambda: owned,
)
await parent.start()
assert owned.is_started is True
await parent.close()
assert owned.is_started is False
asyncio.run(run())
# -- context-bound resolution -------------------------------------------------
def test_resolve_from_context():
async def run():
from reme4.components.application_context import ApplicationContext
target = DepTarget(name="real_index")
ctx = ApplicationContext()
ctx.components = {ComponentEnum.KEYWORD_INDEX: {"real_index": target}}
comp = StubComponent(app_context=ctx)
comp.dep = BaseComponent.bind("real_index", DepTarget)
await comp.start()
assert comp.dep is target
await comp.close()
asyncio.run(run())
def test_resolve_from_context_optional_missing():
async def run():
from reme4.components.application_context import ApplicationContext
ctx = ApplicationContext()
ctx.components = {}
comp = StubComponent(app_context=ctx)
comp.dep = BaseComponent.bind("missing", DepTarget, optional=True)
await comp.start()
assert comp.dep is None
await comp.close()
asyncio.run(run())
def test_resolve_from_context_required_missing_raises():
async def run():
from reme4.components.application_context import ApplicationContext
ctx = ApplicationContext()
ctx.components = {}
comp = StubComponent(app_context=ctx)
comp.dep = BaseComponent.bind("missing", RequiredDepTarget, optional=False)
with pytest.raises(ValueError, match="not found"):
await comp.start()
asyncio.run(run())
# -- ComponentMixin paths -----------------------------------------------------
def test_vault_path_no_context():
mixin = ComponentMixin()
from pathlib import Path
assert mixin.vault_path == Path.cwd()
def test_to_vault_relative_inside_vault():
with tempfile.TemporaryDirectory() as tmpdir:
old_cwd = os.getcwd()
os.chdir(tmpdir)
try:
mixin = ComponentMixin()
abs_path = mixin.vault_path / "sub" / "file.md"
rel = mixin.to_vault_relative(abs_path)
assert rel == str(abs_path.relative_to(mixin.vault_path))
finally:
os.chdir(old_cwd)
def test_to_vault_relative_outside_vault():
mixin = ComponentMixin()
result = mixin.to_vault_relative("/some/other/path")
assert result == "/some/other/path"
# -- vault metadata paths -----------------------------------------------------
def test_vault_metadata_path_no_context():
comp = StubComponent()
assert comp.vault_metadata_path.name == "metadata"
def test_component_metadata_path():
comp = StubComponent()
assert comp.component_metadata_path.name == ComponentEnum.FILE_PARSER.value
if __name__ == "__main__":
print("\n=== BaseComponent Tests ===")
test_dependency_repr_optional()
test_dependency_repr_required()
test_dependency_getattr_raises()
test_bind_returns_none_for_empty_name()
test_bind_returns_dependency_placeholder()
test_bind_rejects_base_component_type()
test_bind_rejects_no_component_type()
test_bind_with_default_factory()
test_bind_optional_flag()
test_dependencies_lists_unresolved()
test_start_close_idempotent()
test_restart()
test_async_context_manager()
test_resolve_standalone_optional_becomes_none()
test_resolve_standalone_with_default_factory()
test_resolve_standalone_required_no_factory_keeps_placeholder()
test_owned_components_started_and_closed()
test_resolve_from_context()
test_resolve_from_context_optional_missing()
test_resolve_from_context_required_missing_raises()
test_vault_path_no_context()
test_to_vault_relative_outside_vault()
test_vault_metadata_path_no_context()
test_component_metadata_path()
print("\n所有测试通过!")

View file

@ -71,7 +71,7 @@ def _node(path: str, links: list[tuple[str, str | None, str | None]] | None = No
async def _make_store(nodes: list[FileNode]) -> LocalFileStore:
"""LocalFileStore seeded with the given graph nodes (no files on disk)."""
store = LocalFileStore(name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_store="")
await store.start()
if nodes:
await store.file_graph.upsert_nodes(nodes)

View file

@ -0,0 +1,151 @@
"""Tests for ComponentRegistry."""
# pylint: disable=missing-function-docstring,missing-class-docstring,protected-access,unused-argument
import pytest
from reme4.components.base_component import BaseComponent
from reme4.components.component_registry import ComponentRegistry
from reme4.enumeration import ComponentEnum
class _DummyComponent(BaseComponent):
component_type = ComponentEnum.FILE_PARSER
class _AnotherComponent(BaseComponent):
component_type = ComponentEnum.KEYWORD_INDEX
class _NoComponentType:
pass
class _BaseComponentType(BaseComponent):
component_type = ComponentEnum.BASE
# -- register & get -----------------------------------------------------------
def test_register_direct_with_explicit_name():
reg = ComponentRegistry()
reg.register(_DummyComponent, "my_parser")
assert reg.get(ComponentEnum.FILE_PARSER, "my_parser") is _DummyComponent
def test_register_direct_defaults_to_class_name():
reg = ComponentRegistry()
reg.register(_DummyComponent)
assert reg.get(ComponentEnum.FILE_PARSER, "_DummyComponent") is _DummyComponent
def test_register_decorator():
reg = ComponentRegistry()
@reg.register("alias")
class MyParser(BaseComponent):
component_type = ComponentEnum.FILE_PARSER
assert reg.get(ComponentEnum.FILE_PARSER, "alias") is MyParser
def test_register_overwrite_warns(caplog):
reg = ComponentRegistry()
reg.register(_DummyComponent, "dup")
reg.register(_DummyComponent, "dup")
assert reg.get(ComponentEnum.FILE_PARSER, "dup") is _DummyComponent
def test_register_rejects_missing_component_type():
reg = ComponentRegistry()
with pytest.raises(TypeError, match="ComponentEnum"):
reg.register(_NoComponentType, "bad")
def test_register_rejects_empty_name():
reg = ComponentRegistry()
with pytest.raises(ValueError, match="empty"):
reg._do_register(_DummyComponent, "")
def test_register_rejects_non_class_non_string():
reg = ComponentRegistry()
with pytest.raises(TypeError, match="Expected a class or string"):
reg.register(42)
# -- get_all ------------------------------------------------------------------
def test_get_all_returns_copy():
reg = ComponentRegistry()
reg.register(_DummyComponent, "a")
reg.register(_AnotherComponent, "b")
parsers = reg.get_all(ComponentEnum.FILE_PARSER)
assert parsers == {"a": _DummyComponent}
indexes = reg.get_all(ComponentEnum.KEYWORD_INDEX)
assert indexes == {"b": _AnotherComponent}
# Mutating the copy doesn't affect the registry.
parsers["hacked"] = _DummyComponent
assert "hacked" not in reg.get_all(ComponentEnum.FILE_PARSER)
def test_get_all_unknown_type_returns_empty():
reg = ComponentRegistry()
assert not reg.get_all(ComponentEnum.LLM)
# -- get (miss) ---------------------------------------------------------------
def test_get_nonexistent_returns_none():
reg = ComponentRegistry()
assert reg.get(ComponentEnum.FILE_PARSER, "nope") is None
# -- unregister ---------------------------------------------------------------
def test_unregister_existing():
reg = ComponentRegistry()
reg.register(_DummyComponent, "x")
assert reg.unregister(ComponentEnum.FILE_PARSER, "x") is True
assert reg.get(ComponentEnum.FILE_PARSER, "x") is None
def test_unregister_missing_returns_false():
reg = ComponentRegistry()
assert reg.unregister(ComponentEnum.FILE_PARSER, "nope") is False
# -- clear --------------------------------------------------------------------
def test_clear():
reg = ComponentRegistry()
reg.register(_DummyComponent, "a")
reg.register(_AnotherComponent, "b")
reg.clear()
assert not reg.get_all(ComponentEnum.FILE_PARSER)
assert not reg.get_all(ComponentEnum.KEYWORD_INDEX)
if __name__ == "__main__":
print("\n=== ComponentRegistry Tests ===")
test_register_direct_with_explicit_name()
test_register_direct_defaults_to_class_name()
test_register_decorator()
test_register_rejects_missing_component_type()
test_register_rejects_empty_name()
test_register_rejects_non_class_non_string()
test_get_all_returns_copy()
test_get_all_unknown_type_returns_empty()
test_get_nonexistent_returns_none()
test_unregister_existing()
test_unregister_missing_returns_false()
test_clear()
print("\n所有测试通过!")

View file

@ -67,7 +67,7 @@ class temp_chdir:
async def _make_store(files: dict[str, str] | None = None) -> LocalFileStore:
"""LocalFileStore seeded with files on disk + registered in the graph."""
store = LocalFileStore(name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_store="")
await store.start()
nodes: list[FileNode] = []
for rel, content in (files or {}).items():

View file

@ -73,7 +73,7 @@ async def _make_store_with_dailies(entries: list[tuple[str, str, str]]) -> Local
``daily/<date>/<session_id>.md`` with a minimal ``name``-only
frontmatter no opinionated status / lifecycle axes.
"""
store = LocalFileStore(name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_store="")
await store.start()
for day, session_id, body in entries:
day_dir = Path.cwd() / "daily" / day
@ -156,7 +156,7 @@ def test_daily_list_returns_path_session_id_metadata():
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = LocalFileStore(name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_store="")
await store.start()
await _seed_note(
"2026-05-18",
@ -212,7 +212,7 @@ def test_daily_list_empty_when_no_daily_dir():
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = LocalFileStore(name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_store="")
await store.start()
step = daily_list_step.DailyListStep(file_store=store)
await step(date="2026-05-18")
@ -439,7 +439,7 @@ def test_day_index_lists_each_note():
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = LocalFileStore(name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_store="")
await store.start()
await _seed_note("2026-05-18", "alpha", name="Alpha Project")
await _seed_note("2026-05-18", "beta", name="Beta Project")
@ -461,7 +461,7 @@ def test_day_index_includes_note_descriptions():
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = LocalFileStore(name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_store="")
await store.start()
cases = [
("alpha", "Alpha Project", "实现 JWT auth 中间件,迁移 session middleware"),
@ -516,7 +516,7 @@ def test_day_index_preserves_user_content_outside_marker():
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = LocalFileStore(name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_store="")
await store.start()
await _seed_note("2026-05-18", "alpha")
reindex = daily_reindex_step.DailyReindexStep(file_store=store)

View file

@ -1,572 +0,0 @@
"""Tests for LocalFileStore and FaissLocalFileStore."""
# pylint: disable=protected-access
import asyncio
import hashlib
import importlib.util
import os
import tempfile
import warnings
import numpy as np
from reme4.components.embedding import BaseEmbeddingModel
from reme4.components.file_store import FaissLocalFileStore, LocalFileStore
from reme4.schema import FileChunk, FileNode
warnings.filterwarnings("ignore", category=DeprecationWarning, module="jieba")
warnings.filterwarnings("ignore", category=DeprecationWarning, module="pkg_resources")
_FAISS_AVAILABLE = importlib.util.find_spec("faiss") is not None
class FakeEmbeddingModel(BaseEmbeddingModel):
"""Deterministic stub: each lowercased word adds 1.0 at hash(word) % dim."""
def __init__(self, dimensions: int = 8, **kwargs):
super().__init__(model_name="fake", dimensions=dimensions, enable_cache=False, **kwargs)
async def _get_embeddings(self, input_text, **kwargs):
out = []
for t in input_text:
v = np.zeros(self.dimensions, dtype=np.float32)
for w in t.lower().split():
idx = int.from_bytes(hashlib.md5(w.encode()).digest()[:2], "big") % self.dimensions
v[idx] += 1.0
out.append(v.tolist())
return out
async def health_check(self, timeout: float = 2.0) -> bool:
self.is_healthy = True
return True
class temp_chdir:
"""Context manager to temporarily chdir into a path and restore on exit."""
def __init__(self, path):
self.path = path
self.old = None
def __enter__(self):
self.old = os.getcwd()
os.chdir(self.path)
return self
def __exit__(self, *exc):
os.chdir(self.old)
async def make_store(store_name: str = "test_store", **kwargs) -> LocalFileStore:
"""Build a started LocalFileStore with embedding disabled (no OpenAI dep)."""
store = LocalFileStore(name=store_name, embedding_model="", **kwargs)
await store.start()
return store
def make_file(
path: str,
text: str,
chunk_count: int = 1,
) -> tuple[FileNode, list[FileChunk]]:
"""Build a (FileNode, [FileChunk]) tuple ready for upsert_file."""
chunks = [
FileChunk(id=f"{path}::chunk{i}", path=path, text=f"{text} part{i}", start_line=i, end_line=i + 1)
for i in range(chunk_count)
]
node = FileNode(path=path, st_mtime=1.0, chunk_ids=[c.id for c in chunks])
return node, chunks
def test_upsert_single_file():
"""upsert_file with a one-element list stores chunks and node."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
store = await make_store()
node, chunks = make_file("a.md", "hello world", chunk_count=2)
await store.upsert([(node, chunks)])
# Chunks landed in memory
assert len(store.file_chunks) == 2
assert {c.path for c in store.file_chunks.values()} == {"a.md"}
# Node landed in graph
nodes = await store.get_nodes(["a.md"])
assert len(nodes) == 1
assert sorted(nodes[0].chunk_ids) == sorted([c.id for c in chunks])
await store.close()
print("✓ test_upsert_single_file passed")
asyncio.run(run())
def test_upsert_multiple_files():
"""upsert_file accepts a list of tuples and indexes them all."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
store = await make_store()
files = [make_file("a.md", "alpha"), make_file("b.md", "beta")]
await store.upsert(files)
assert len(store.file_chunks) == 2
paths = {n.path for n in await store.get_nodes()}
assert paths == {"a.md", "b.md"}
await store.close()
print("✓ test_upsert_multiple_files passed")
asyncio.run(run())
def test_upsert_replaces_old_chunks():
"""Re-upserting the same path points the node at the new chunk set."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
store = await make_store()
n1, c1 = make_file("a.md", "v1", chunk_count=2)
await store.upsert([(n1, c1)])
# Different chunks for the same path
n2 = FileNode(path="a.md", st_mtime=2.0)
c2 = [FileChunk(id="a.md::new", path="a.md", text="v2 only", start_line=0, end_line=1)]
n2.chunk_ids = [c.id for c in c2]
await store.upsert([(n2, c2)])
# The node now references the new chunk set, not the old one.
nodes = await store.get_nodes(["a.md"])
assert nodes[0].chunk_ids == ["a.md::new"]
assert "a.md::new" in store.file_chunks
await store.close()
print("✓ test_upsert_replaces_old_chunks passed")
asyncio.run(run())
def test_delete_by_path_single():
"""delete_by_path drops chunks and the node entry."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
store = await make_store()
await store.upsert([make_file("a.md", "alpha"), make_file("b.md", "beta")])
await store.delete("a.md")
assert all(c.path != "a.md" for c in store.file_chunks.values())
assert {n.path for n in await store.get_nodes()} == {"b.md"}
await store.close()
print("✓ test_delete_by_path_single passed")
asyncio.run(run())
def test_delete_by_path_list():
"""delete_by_path accepts a list of paths."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
store = await make_store()
await store.upsert(
[
make_file("a.md", "alpha"),
make_file("b.md", "beta"),
make_file("c.md", "gamma"),
],
)
await store.delete(["a.md", "b.md"])
assert {n.path for n in await store.get_nodes()} == {"c.md"}
assert all(c.path == "c.md" for c in store.file_chunks.values())
await store.close()
print("✓ test_delete_by_path_list passed")
asyncio.run(run())
def test_delete_by_path_missing_is_noop():
"""Deleting a nonexistent path is a no-op."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
store = await make_store()
await store.upsert([make_file("a.md", "alpha")])
before = len(store.file_chunks)
await store.delete("ghost.md")
assert len(store.file_chunks) == before
await store.close()
print("✓ test_delete_by_path_missing_is_noop passed")
asyncio.run(run())
def test_clear():
"""clear() empties chunks and the file graph."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
store = await make_store()
await store.upsert([make_file("a.md", "alpha"), make_file("b.md", "beta")])
await store.clear()
assert store.file_chunks == {}
assert await store.get_nodes() == []
await store.close()
print("✓ test_clear passed")
asyncio.run(run())
def test_keyword_search():
"""keyword_search returns matching chunks ranked by BM25 score."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
store = await make_store()
await store.upsert(
[
make_file("a.md", "python programming language"),
make_file("b.md", "java programming language"),
make_file("c.md", "python data analysis"),
],
)
results = await store.keyword_search("python", limit=5, search_filter={})
paths = {r.path for r in results}
assert "a.md" in paths or "c.md" in paths
# Each result should carry a keyword score.
for r in results:
assert r.scores.get("keyword", 0) > 0
await store.close()
print("✓ test_keyword_search passed")
asyncio.run(run())
def test_keyword_search_empty_query():
"""Empty/whitespace queries return no results."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
store = await make_store()
await store.upsert([make_file("a.md", "hello")])
assert await store.keyword_search("", limit=5, search_filter={}) == []
assert await store.keyword_search(" ", limit=5, search_filter={}) == []
await store.close()
print("✓ test_keyword_search_empty_query passed")
asyncio.run(run())
def test_vector_search_disabled_returns_empty():
"""Without an embedding model, vector_search returns []."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
store = await make_store()
await store.upsert([make_file("a.md", "hello")])
assert store.embedding_model is None
assert await store.vector_search("hello", limit=5, search_filter={}) == []
await store.close()
print("✓ test_vector_search_disabled_returns_empty passed")
asyncio.run(run())
def test_persistence_roundtrip():
"""close() dumps chunks; a fresh store loads them from disk."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
s1 = await make_store()
await s1.upsert([make_file("a.md", "alpha"), make_file("b.md", "beta")])
await s1.close()
s2 = await make_store()
assert {c.path for c in s2.file_chunks.values()} == {"a.md", "b.md"}
# Graph should also be persisted independently via its own dump.
assert {n.path for n in await s2.get_nodes()} == {"a.md", "b.md"}
await s2.close()
print("✓ test_persistence_roundtrip passed")
asyncio.run(run())
def test_rebuild_links_delegates_to_graph():
"""rebuild_links() on the store delegates to the underlying file_graph."""
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
store = await make_store()
from reme4.schema import FileLink
node = FileNode(
path="a.md",
st_mtime=1.0,
links=[FileLink(source_path="a.md", target_path="b.md")],
)
chunks = [FileChunk(id="a::1", path="a.md", text="x", start_line=0, end_line=1)]
node.chunk_ids = [c.id for c in chunks]
await store.upsert([(node, chunks), make_file("b.md", "beta")])
await store.rebuild_links()
inlinks = await store.get_inlinks("b.md")
assert {lnk.source_path for lnk in inlinks} == {"a.md"}
await store.close()
print("✓ test_rebuild_links_delegates_to_graph passed")
asyncio.run(run())
# --- FaissLocalFileStore tests --------------------------------------------------
def _skip_if_no_faiss(name: str) -> bool:
if not _FAISS_AVAILABLE:
print(f"{name} skipped (faiss not installed)")
return True
return False
async def make_faiss_store(store_name: str = "test_faiss", **kwargs) -> FaissLocalFileStore:
"""Build a started FaissLocalFileStore wired to FakeEmbeddingModel (no API calls)."""
store = FaissLocalFileStore(name=store_name, embedding_model="fake", **kwargs)
fake = FakeEmbeddingModel()
# Replace the unresolved Dependency placeholder with a concrete instance and
# let start() cascade lifecycle to it via _owned.
store.embedding_model = fake
store._owned.append(fake)
await store.start()
return store
def test_faiss_vector_search_basic():
"""vector_search returns chunks ranked by cosine similarity to the query."""
if _skip_if_no_faiss("test_faiss_vector_search_basic"):
return
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
store = await make_faiss_store()
await store.upsert(
[
make_file("a.md", "alpha"),
make_file("b.md", "beta"),
make_file("c.md", "alpha gamma"),
],
)
results = await store.vector_search("alpha", limit=3, search_filter={})
assert results, "vector_search returned no results"
for r in results:
assert "vector" in r.scores
assert r.scores["score"] == r.scores["vector"]
# Top hit should match an "alpha"-bearing doc.
assert results[0].path in {"a.md", "c.md"}
# All distinct chunks (no duplicates from tombstones).
assert len({r.id for r in results}) == len(results)
await store.close()
print("✓ test_faiss_vector_search_basic passed")
asyncio.run(run())
def test_faiss_persistence_roundtrip():
"""close() writes FAISS sidecar; a fresh store loads it without rebuilding."""
if _skip_if_no_faiss("test_faiss_persistence_roundtrip"):
return
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
s1 = await make_faiss_store()
await s1.upsert([make_file("a.md", "alpha"), make_file("b.md", "beta")])
r1 = await s1.vector_search("alpha", limit=2, search_filter={})
assert s1.faiss_path.exists() is False # not yet dumped
await s1.close()
assert s1.faiss_path.exists() and s1.faiss_idmap_path.exists()
s2 = await make_faiss_store()
assert s2._faiss_index is not None
assert s2._faiss_index.ntotal == 2
r2 = await s2.vector_search("alpha", limit=2, search_filter={})
assert [r.path for r in r2] == [r.path for r in r1]
await s2.close()
print("✓ test_faiss_persistence_roundtrip passed")
asyncio.run(run())
def test_faiss_delete_removes_from_search():
"""Deleting a file tombstones its chunks; subsequent search excludes them."""
if _skip_if_no_faiss("test_faiss_delete_removes_from_search"):
return
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
store = await make_faiss_store()
await store.upsert([make_file("a.md", "alpha"), make_file("b.md", "alpha beta")])
await store.delete("a.md")
results = await store.vector_search("alpha", limit=5, search_filter={})
assert all(r.path != "a.md" for r in results)
# All tombstoned rows still present in id_map; live mapping shrank.
assert len(store._id_to_row) == 1
assert len(store._tombstones) == 1
await store.close()
print("✓ test_faiss_delete_removes_from_search passed")
asyncio.run(run())
def test_faiss_upsert_replaces_vectors():
"""Re-upserting a path with new chunk ids tombstones the old vectors."""
if _skip_if_no_faiss("test_faiss_upsert_replaces_vectors"):
return
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
store = await make_faiss_store()
n1, c1 = make_file("a.md", "alpha", chunk_count=2)
await store.upsert([(n1, c1)])
assert store._faiss_index.ntotal == 2
assert len(store._tombstones) == 0
# New chunk ids for the same path → old ones become tombstones.
n2 = FileNode(path="a.md", st_mtime=2.0)
c2 = [FileChunk(id="a.md::new", path="a.md", text="gamma", start_line=0, end_line=1)]
n2.chunk_ids = [c.id for c in c2]
await store.upsert([(n2, c2)])
assert store._faiss_index.ntotal == 3 # 2 old + 1 new appended
assert len(store._tombstones) == 2 # both old rows tombstoned
assert "a.md::new" in store._id_to_row
results = await store.vector_search("gamma", limit=5, search_filter={})
assert results and results[0].id == "a.md::new"
await store.close()
print("✓ test_faiss_upsert_replaces_vectors passed")
asyncio.run(run())
def test_faiss_clear_empties_index():
"""clear() resets FAISS state and removes sidecar files."""
if _skip_if_no_faiss("test_faiss_clear_empties_index"):
return
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
store = await make_faiss_store()
await store.upsert([make_file("a.md", "alpha")])
await store.dump()
assert store.faiss_path.exists()
await store.clear()
assert store._faiss_index.ntotal == 0
assert store._id_map == [] and store._id_to_row == {}
assert not store.faiss_path.exists()
assert not store.faiss_idmap_path.exists()
assert await store.vector_search("alpha", limit=5, search_filter={}) == []
await store.close()
print("✓ test_faiss_clear_empties_index passed")
asyncio.run(run())
def test_faiss_rebuild_when_sidecar_missing():
"""If the FAISS sidecar is missing on load, the index rebuilds from chunks JSONL."""
if _skip_if_no_faiss("test_faiss_rebuild_when_sidecar_missing"):
return
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
s1 = await make_faiss_store()
await s1.upsert([make_file("a.md", "alpha"), make_file("b.md", "beta")])
await s1.close()
# Drop the FAISS sidecar but keep chunks JSONL — load() must rebuild.
s1.faiss_path.unlink()
s1.faiss_idmap_path.unlink()
s2 = await make_faiss_store()
assert s2._faiss_index is not None and s2._faiss_index.ntotal == 2
results = await s2.vector_search("alpha", limit=2, search_filter={})
assert any(r.path == "a.md" for r in results)
await s2.close()
print("✓ test_faiss_rebuild_when_sidecar_missing passed")
asyncio.run(run())
def test_faiss_disabled_without_embedding():
"""embedding_model="" → FAISS path stays dormant; vector_search returns []."""
if _skip_if_no_faiss("test_faiss_disabled_without_embedding"):
return
async def run():
with tempfile.TemporaryDirectory() as tmpdir, temp_chdir(tmpdir):
store = FaissLocalFileStore(name="disabled", embedding_model="")
await store.start()
await store.upsert([make_file("a.md", "alpha")])
assert store.embedding_model is None
assert store._faiss_index is None
assert await store.vector_search("alpha", limit=5, search_filter={}) == []
await store.close()
print("✓ test_faiss_disabled_without_embedding passed")
asyncio.run(run())
if __name__ == "__main__":
print("\n=== LocalFileStore Tests ===")
test_upsert_single_file()
test_upsert_multiple_files()
test_upsert_replaces_old_chunks()
test_delete_by_path_single()
test_delete_by_path_list()
test_delete_by_path_missing_is_noop()
test_clear()
test_keyword_search()
test_keyword_search_empty_query()
test_vector_search_disabled_returns_empty()
test_persistence_roundtrip()
test_rebuild_links_delegates_to_graph()
print("\n=== FaissLocalFileStore Tests ===")
test_faiss_vector_search_basic()
test_faiss_persistence_roundtrip()
test_faiss_delete_removes_from_search()
test_faiss_upsert_replaces_vectors()
test_faiss_clear_empties_index()
test_faiss_rebuild_when_sidecar_missing()
test_faiss_disabled_without_embedding()
print("\n所有测试通过!")

247
tests4/unit/test_job.py Normal file
View file

@ -0,0 +1,247 @@
"""Tests for BaseJob and BackgroundJob."""
# pylint: disable=protected-access,missing-function-docstring,missing-class-docstring,no-self-argument,unused-argument
import asyncio
from unittest.mock import AsyncMock, MagicMock
import pytest
from reme4.components.component_registry import ComponentRegistry
from reme4.components.job.background_job import BackgroundJob
from reme4.components.job.base_job import BaseJob
from reme4.schema import ComponentConfig
# -- helpers ------------------------------------------------------------------
def _make_registry_and_context(step_classes=None):
"""Build a fresh registry + minimal app_context for job tests."""
reg = ComponentRegistry()
if step_classes:
for name, cls in step_classes.items():
reg.register(cls, name)
ctx = MagicMock()
ctx.components = {}
return reg, ctx
# -- BaseJob._resolve_step ---------------------------------------------------
def test_resolve_step_missing_backend():
job = BaseJob(name="j")
job.app_context = MagicMock()
with pytest.raises(ValueError, match="missing the required 'backend'"):
job._resolve_step(ComponentConfig(backend=""))
def test_resolve_step_unregistered_backend():
job = BaseJob(name="j")
job.app_context = MagicMock()
with pytest.raises(ValueError, match="Unregistered backend"):
job._resolve_step(ComponentConfig(backend="nonexistent_step"))
# -- BaseJob.__call__ error capture ------------------------------------------
def test_call_captures_exception():
async def run():
failing_step = AsyncMock(side_effect=RuntimeError("boom"))
job = BaseJob(name="j")
job.app_context = MagicMock()
job.step_specs = []
job._build_steps = lambda: [failing_step]
response = await job()
assert response.success is False
assert "boom" in response.answer
asyncio.run(run())
def test_call_runs_steps_in_order():
async def run():
call_order = []
async def step1(ctx):
call_order.append("s1")
async def step2(ctx):
call_order.append("s2")
job = BaseJob(name="j")
job.app_context = MagicMock()
job.step_specs = []
job._build_steps = lambda: [step1, step2]
response = await job()
assert response.success is True
assert call_order == ["s1", "s2"]
asyncio.run(run())
# -- BaseJob._start requires app_context ------------------------------------
def test_start_without_app_context_raises():
async def run():
job = BaseJob(name="j")
with pytest.raises(RuntimeError, match="app_context must be provided"):
await job._start()
asyncio.run(run())
# -- BackgroundJob._backoff_delay --------------------------------------------
def test_backoff_delay_increases():
job = BackgroundJob(
name="bg",
backoff_base=1.0,
backoff_cap=60.0,
)
delays = [job._backoff_delay(i) for i in range(10)]
# Delay should generally increase (with jitter, so we check trend).
assert delays[-1] >= delays[0] or delays[-1] == job.backoff_cap
def test_backoff_delay_capped():
job = BackgroundJob(
name="bg",
backoff_base=1.0,
backoff_cap=10.0,
)
for _ in range(100):
delay = job._backoff_delay(20)
assert delay <= job.backoff_cap
def test_backoff_delay_has_jitter():
job = BackgroundJob(
name="bg",
backoff_base=1.0,
backoff_cap=60.0,
)
delays = {job._backoff_delay(5) for _ in range(20)}
assert len(delays) > 1
def test_backoff_delay_attempt_zero():
job = BackgroundJob(
name="bg",
backoff_base=2.0,
backoff_cap=60.0,
)
for _ in range(50):
delay = job._backoff_delay(0)
assert 0 < delay <= 2.0 * 1.5
# -- BackgroundJob supervisor loop -------------------------------------------
def test_supervisor_restarts_on_crash():
async def run():
call_count = 0
stop = asyncio.Event()
class CrashingJob(BackgroundJob):
async def __call__(self_, **kwargs):
nonlocal call_count
call_count += 1
if call_count < 3:
raise RuntimeError("crash")
stop.set()
job = CrashingJob(
name="bg",
supervisor=True,
backoff_base=0.01,
backoff_cap=0.05,
)
job._stop_event = stop
await job._run_with_supervisor()
assert call_count == 3
asyncio.run(run())
def test_supervisor_disabled_propagates_exception():
async def run():
class FatalJob(BackgroundJob):
async def __call__(self_, **kwargs):
raise RuntimeError("fatal")
job = FatalJob(
name="bg",
supervisor=False,
)
job._stop_event = asyncio.Event()
with pytest.raises(RuntimeError, match="fatal"):
await job._run_with_supervisor()
asyncio.run(run())
# -- BackgroundJob._wait_or_stop ----------------------------------------------
def test_wait_or_stop_returns_on_stop():
async def run():
job = BackgroundJob(name="bg")
job._stop_event = asyncio.Event()
job._stop_event.set()
await job._wait_or_stop(10.0)
asyncio.run(run())
# -- BackgroundJob._shutdown_task ---------------------------------------------
def test_shutdown_task_none():
async def run():
job = BackgroundJob(name="bg")
job._task = None
await job._shutdown_task()
asyncio.run(run())
def test_shutdown_task_cancels_on_timeout():
async def run():
async def hang_forever():
await asyncio.sleep(999)
job = BackgroundJob(name="bg", close_timeout=0.05)
job._task = asyncio.create_task(hang_forever())
await job._shutdown_task()
assert job._task is None
asyncio.run(run())
if __name__ == "__main__":
print("\n=== Job Tests ===")
test_resolve_step_missing_backend()
test_resolve_step_unregistered_backend()
test_call_captures_exception()
test_call_runs_steps_in_order()
test_start_without_app_context_raises()
test_backoff_delay_increases()
test_backoff_delay_capped()
test_backoff_delay_has_jitter()
test_backoff_delay_attempt_zero()
test_supervisor_restarts_on_crash()
test_supervisor_disabled_propagates_exception()
test_wait_or_stop_returns_on_stop()
test_shutdown_task_none()
test_shutdown_task_cancels_on_timeout()
print("\n所有测试通过!")

View file

@ -50,7 +50,7 @@ async def _store_with(files: dict[str, dict]) -> LocalFileStore:
``description`` populate FileFrontMatter so neighbor meta lookups
have something to surface.
"""
store = LocalFileStore(name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_store="")
await store.start()
nodes: list[FileNode] = []
root = Path.cwd()
@ -84,7 +84,7 @@ def test_expand_links_empty_paths_short_circuits():
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
store = LocalFileStore(name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_store="")
await store.start()
result = await expand_links(store, [])
assert result == {}

View file

@ -0,0 +1,304 @@
"""Tests for PromptHandler."""
# pylint: disable=missing-function-docstring
import json
import tempfile
from pathlib import Path
import pytest
import yaml
from reme4.components.prompt_handler import PromptHandler
# -- init & load_prompt_dict --------------------------------------------------
def test_init_filters_non_string_values():
ph = PromptHandler(greeting="hello", count=42, items=[1, 2])
assert ph.data == {"greeting": "hello"}
def test_load_prompt_dict_basic():
ph = PromptHandler()
ph.load_prompt_dict({"a": "alpha", "b": "beta"})
assert ph.data == {"a": "alpha", "b": "beta"}
def test_load_prompt_dict_skips_non_string_values():
ph = PromptHandler()
ph.load_prompt_dict({"good": "ok", "bad": 123})
assert ph.data == {"good": "ok"}
def test_load_prompt_dict_overwrite_true():
ph = PromptHandler(a="old")
ph.load_prompt_dict({"a": "new"}, overwrite=True)
assert ph.data["a"] == "new"
def test_load_prompt_dict_overwrite_false():
ph = PromptHandler(a="old")
ph.load_prompt_dict({"a": "new"}, overwrite=False)
assert ph.data["a"] == "old"
def test_load_prompt_dict_none():
ph = PromptHandler(a="old")
result = ph.load_prompt_dict(None)
assert result is ph
assert ph.data == {"a": "old"}
def test_load_prompt_dict_non_dict():
ph = PromptHandler()
result = ph.load_prompt_dict("not a dict")
assert result is ph
assert ph.data == {}
# -- load from file -----------------------------------------------------------
def test_load_prompt_by_file_yaml():
data = {"greeting": "Hello {name}", "farewell": "Goodbye"}
with tempfile.NamedTemporaryFile(suffix=".yaml", mode="w", delete=False) as f:
yaml.dump(data, f)
f.flush()
ph = PromptHandler()
ph.load_prompt_by_file(f.name)
assert ph.data["greeting"] == "Hello {name}"
assert ph.data["farewell"] == "Goodbye"
Path(f.name).unlink()
def test_load_prompt_by_file_json():
data = {"q1": "What is {topic}?"}
with tempfile.NamedTemporaryFile(suffix=".json", mode="w", delete=False) as f:
json.dump(data, f)
f.flush()
ph = PromptHandler()
ph.load_prompt_by_file(f.name)
assert ph.data["q1"] == "What is {topic}?"
Path(f.name).unlink()
def test_load_prompt_by_file_none():
ph = PromptHandler()
result = ph.load_prompt_by_file(None)
assert result is ph
def test_load_prompt_by_file_nonexistent():
ph = PromptHandler()
result = ph.load_prompt_by_file("/nonexistent/path.yaml")
assert result is ph
assert ph.data == {}
def test_load_prompt_by_file_unsupported_extension():
with tempfile.NamedTemporaryFile(suffix=".txt", mode="w", delete=False) as f:
f.write("hello")
f.flush()
ph = PromptHandler()
result = ph.load_prompt_by_file(f.name)
assert result is ph
assert ph.data == {}
Path(f.name).unlink()
# -- get_prompt & i18n --------------------------------------------------------
def test_get_prompt_bare_key():
ph = PromptHandler(greeting="Hello")
assert ph.get_prompt("greeting") == "Hello"
def test_get_prompt_strips():
ph = PromptHandler(greeting=" Hello \n")
assert ph.get_prompt("greeting") == "Hello"
def test_get_prompt_missing_raises():
ph = PromptHandler()
with pytest.raises(KeyError, match="not found"):
ph.get_prompt("missing")
def test_get_prompt_language_fallback():
ph = PromptHandler(language="zh", greeting="Hello", greeting_zh="你好")
assert ph.get_prompt("greeting") == "你好"
def test_get_prompt_language_fallback_to_bare():
ph = PromptHandler(language="zh", greeting="Hello")
assert ph.get_prompt("greeting") == "Hello"
def test_has_prompt():
ph = PromptHandler(greeting="Hello")
assert ph.has_prompt("greeting") is True
assert ph.has_prompt("missing") is False
def test_has_prompt_with_language():
ph = PromptHandler(language="en", greeting_en="Hi")
assert ph.has_prompt("greeting") is True
# -- list_prompts -------------------------------------------------------------
def test_list_prompts_all():
ph = PromptHandler(a="1", b_en="2", c_zh="3")
assert sorted(ph.list_prompts()) == ["a", "b_en", "c_zh"]
def test_list_prompts_filtered():
ph = PromptHandler(a="1", b_en="2", c_en="3", d_zh="4")
assert sorted(ph.list_prompts("en")) == ["b_en", "c_en"]
# -- prompt_format (flag filtering) -------------------------------------------
def test_flag_filter_keeps_matching():
ph = PromptHandler()
ph.load_prompt_dict({"p": "[verbose] extra detail\nalways here"})
result = ph.prompt_format("p", verbose=True)
assert "extra detail" in result
assert "always here" in result
def test_flag_filter_removes_non_matching():
ph = PromptHandler()
ph.load_prompt_dict({"p": "[verbose] extra detail\nalways here"})
result = ph.prompt_format("p", verbose=False)
assert "extra detail" not in result
assert "always here" in result
def test_flag_filter_default_false():
ph = PromptHandler()
ph.load_prompt_dict({"p": "[debug] debug info\nbase"})
# When no flags are passed at all, _apply_flag_filter is not called,
# so flagged lines are kept as-is (including the tag text after regex sub).
result = ph.prompt_format("p", debug=False)
assert "debug info" not in result
assert "base" in result
def test_flag_filter_unflagged_lines_always_kept():
ph = PromptHandler()
ph.load_prompt_dict({"p": "line1\nline2\nline3"})
result = ph.prompt_format("p")
assert "line1" in result
assert "line2" in result
assert "line3" in result
# -- prompt_format (variable substitution) ------------------------------------
def test_format_variables():
ph = PromptHandler()
ph.load_prompt_dict({"p": "Hello {name}, welcome to {place}"})
result = ph.prompt_format("p", name="Alice", place="Wonderland")
assert result == "Hello Alice, welcome to Wonderland"
def test_format_missing_variable_raises():
ph = PromptHandler()
ph.load_prompt_dict({"p": "Hello {name}"})
with pytest.raises(ValueError, match="Missing format variables"):
ph.prompt_format("p")
def test_format_missing_variable_no_validate():
ph = PromptHandler()
ph.load_prompt_dict({"p": "Hello {name}"})
result = ph.prompt_format("p", validate=False)
assert "{name}" in result
def test_format_no_variables_no_error():
ph = PromptHandler()
ph.load_prompt_dict({"p": "No vars here"})
result = ph.prompt_format("p")
assert result == "No vars here"
# -- prompt_format (combined flags + variables) -------------------------------
def test_format_flags_and_variables_combined():
ph = PromptHandler()
ph.load_prompt_dict(
{
"p": "[verbose] Debug: {detail}\nResult: {answer}",
},
)
result = ph.prompt_format("p", verbose=True, detail="trace", answer="42")
assert "Debug: trace" in result
assert "Result: 42" in result
def test_format_flags_false_variable_not_needed():
ph = PromptHandler()
ph.load_prompt_dict(
{
"p": "[verbose] Debug: {detail}\nResult: {answer}",
},
)
result = ph.prompt_format("p", verbose=False, answer="42")
assert "Debug" not in result
assert "Result: 42" in result
# -- repr ---------------------------------------------------------------------
def test_repr():
ph = PromptHandler(language="en", a="1", b="2")
r = repr(ph)
assert "en" in r
assert "2" in r
if __name__ == "__main__":
print("\n=== PromptHandler Tests ===")
test_init_filters_non_string_values()
test_load_prompt_dict_basic()
test_load_prompt_dict_skips_non_string_values()
test_load_prompt_dict_overwrite_true()
test_load_prompt_dict_overwrite_false()
test_load_prompt_dict_none()
test_load_prompt_dict_non_dict()
test_load_prompt_by_file_yaml()
test_load_prompt_by_file_json()
test_load_prompt_by_file_none()
test_load_prompt_by_file_nonexistent()
test_load_prompt_by_file_unsupported_extension()
test_get_prompt_bare_key()
test_get_prompt_strips()
test_get_prompt_missing_raises()
test_get_prompt_language_fallback()
test_get_prompt_language_fallback_to_bare()
test_has_prompt()
test_has_prompt_with_language()
test_list_prompts_all()
test_list_prompts_filtered()
test_flag_filter_keeps_matching()
test_flag_filter_removes_non_matching()
test_flag_filter_default_false()
test_flag_filter_unflagged_lines_always_kept()
test_format_variables()
test_format_missing_variable_raises()
test_format_missing_variable_no_validate()
test_format_no_variables_no_error()
test_format_flags_and_variables_combined()
test_format_flags_false_variable_not_needed()
test_repr()
print("\n所有测试通过!")

View file

@ -48,7 +48,7 @@ def _run(coro):
async def _make_store() -> LocalFileStore:
store = LocalFileStore(name="t_img", embedding_model="")
store = LocalFileStore(name="t_img", embedding_store="")
await store.start()
return store

View file

@ -50,7 +50,7 @@ def _run(coro):
async def _store_with(files: dict[str, dict]) -> LocalFileStore:
"""LocalFileStore seeded with files + parsed wikilinks + optional frontmatter."""
store = LocalFileStore(name="t_read_neighbors", embedding_model="")
store = LocalFileStore(name="t_read_neighbors", embedding_store="")
await store.start()
nodes: list[FileNode] = []
root = Path.cwd()
@ -169,7 +169,7 @@ def test_read_with_neighbors_non_md_falls_through():
async def run():
with tempfile.TemporaryDirectory() as tmp, temp_chdir(tmp):
(Path(tmp) / "notes.txt").write_text("plain text body", encoding="utf-8")
store = LocalFileStore(name="t_read_neighbors_nonmd", embedding_model="")
store = LocalFileStore(name="t_read_neighbors_nonmd", embedding_store="")
await store.start()
resp = await _read(store, step_kwargs={"with_neighbors": True}, path="notes.txt")
assert resp.success is True

View file

@ -49,7 +49,7 @@ class temp_chdir:
async def _make_store() -> LocalFileStore:
"""Minimal LocalFileStore (embedding disabled). vault_path resolves to CWD."""
store = LocalFileStore(name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_store="")
await store.start()
return store

View file

@ -0,0 +1,181 @@
"""Tests for RuntimeContext."""
# pylint: disable=protected-access,missing-function-docstring
import asyncio
import pytest
from reme4.components.runtime_context import RuntimeContext
from reme4.enumeration import ChunkEnum
# -- dict-like access ---------------------------------------------------------
def test_getitem_setitem():
ctx = RuntimeContext(foo="bar")
assert ctx["foo"] == "bar"
ctx["baz"] = 42
assert ctx["baz"] == 42
def test_getitem_missing_raises():
ctx = RuntimeContext()
with pytest.raises(KeyError):
_ = ctx["nope"]
def test_contains():
ctx = RuntimeContext(a=1)
assert "a" in ctx
assert "b" not in ctx
def test_delitem():
ctx = RuntimeContext(a=1)
del ctx["a"]
assert "a" not in ctx
def test_get_with_default():
ctx = RuntimeContext(a=1)
assert ctx.get("a") == 1
assert ctx.get("b", "fallback") == "fallback"
assert ctx.get("b") is None
def test_update_merges_and_returns_self():
ctx = RuntimeContext(a=1)
result = ctx.update({"b": 2, "c": 3})
assert result is ctx
assert ctx["b"] == 2
assert ctx["c"] == 3
# -- from_context -------------------------------------------------------------
def test_from_context_creates_new_when_none():
ctx = RuntimeContext.from_context(None, x=10)
assert ctx["x"] == 10
def test_from_context_reuses_existing():
original = RuntimeContext(a=1)
reused = RuntimeContext.from_context(original, b=2)
assert reused is original
assert reused["a"] == 1
assert reused["b"] == 2
# -- apply_mapping ------------------------------------------------------------
def test_apply_mapping_copies_values():
ctx = RuntimeContext(src="hello")
result = ctx.apply_mapping({"src": "dst"})
assert result is ctx
assert ctx["dst"] == "hello"
assert ctx["src"] == "hello"
def test_apply_mapping_skips_missing_source():
ctx = RuntimeContext(a=1)
ctx.apply_mapping({"missing_key": "target"})
assert "target" not in ctx
def test_apply_mapping_empty_is_noop():
ctx = RuntimeContext(a=1)
result = ctx.apply_mapping({})
assert result is ctx
# -- streaming ----------------------------------------------------------------
def test_stream_property():
ctx_no_queue = RuntimeContext()
assert ctx_no_queue.stream is False
ctx_with_queue = RuntimeContext(stream_queue=asyncio.Queue())
assert ctx_with_queue.stream is True
def test_enqueue_raises_without_queue():
async def run():
ctx = RuntimeContext()
with pytest.raises(RuntimeError, match="Stream queue not initialized"):
await ctx._enqueue(None)
asyncio.run(run())
def test_add_stream_string():
async def run():
q = asyncio.Queue()
ctx = RuntimeContext(stream_queue=q)
result = await ctx.add_stream_string("hello", ChunkEnum.CONTENT)
assert result is ctx
chunk = q.get_nowait()
assert chunk.chunk == "hello"
assert chunk.chunk_type == ChunkEnum.CONTENT
assert chunk.done is False
asyncio.run(run())
def test_add_stream_done():
async def run():
q = asyncio.Queue()
ctx = RuntimeContext(stream_queue=q)
result = await ctx.add_stream_done()
assert result is ctx
chunk = q.get_nowait()
assert chunk.chunk_type == ChunkEnum.DONE
assert chunk.done is True
asyncio.run(run())
# -- response -----------------------------------------------------------------
def test_default_response():
ctx = RuntimeContext()
assert ctx.response.success is True
assert ctx.response.answer == ""
def test_custom_response():
from reme4.schema import Response
resp = Response(answer="ok", success=False)
ctx = RuntimeContext(response=resp)
assert ctx.response is resp
assert ctx.response.success is False
if __name__ == "__main__":
print("\n=== RuntimeContext Tests ===")
test_getitem_setitem()
test_getitem_missing_raises()
test_contains()
test_delitem()
test_get_with_default()
test_update_merges_and_returns_self()
test_from_context_creates_new_when_none()
test_from_context_reuses_existing()
test_apply_mapping_copies_values()
test_apply_mapping_skips_missing_source()
test_apply_mapping_empty_is_noop()
test_stream_property()
test_enqueue_raises_without_queue()
test_add_stream_string()
test_add_stream_done()
test_default_response()
test_custom_response()
print("\n所有测试通过!")

View file

@ -51,7 +51,7 @@ async def _store_with(files: dict[str, str]) -> LocalFileStore:
Without the parsed links the reverse-index lookup yields nothing and
retarget becomes a no-op.
"""
store = LocalFileStore(name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_store="")
await store.start()
nodes: list[FileNode] = []
root = Path.cwd()
@ -72,7 +72,7 @@ async def _store_with(files: dict[str, str]) -> LocalFileStore:
async def _empty_store() -> LocalFileStore:
store = LocalFileStore(name="t", embedding_model="")
store = LocalFileStore(name="t", embedding_store="")
await store.start()
return store

View file

@ -50,7 +50,7 @@ def _run(coro):
async def _make_store() -> LocalFileStore:
store = LocalFileStore(name="t_write_meta", embedding_model="")
store = LocalFileStore(name="t_write_meta", embedding_store="")
await store.start()
return store