From 4348148b723fb9fd30d3e9fe17acf8147996d3a0 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Thu, 22 Jan 2026 16:25:20 +0800 Subject: [PATCH] refactor(core): restructure core modules and update pre-commit configuration --- .pre-commit-config.yaml | 13 +- {reme_ai/core => reme}/config/__init__.py | 0 {reme_ai/core => reme}/config/default.yaml | 5 +- .../config/reme_config_parser.py | 2 +- reme/core/context/__init__.py | 4 + .../core/context/runtime_context.py | 13 +- reme/core/context/service_context.py | 230 ++++++++ {reme_ai => reme}/core/embedding/__init__.py | 4 + .../core/embedding/base_embedding_model.py | 0 .../core/embedding/openai_embedding_model.py | 2 - .../embedding/openai_embedding_model_sync.py | 2 - {reme_ai => reme}/core/flow/__init__.py | 2 - {reme_ai => reme}/core/flow/base_flow.py | 67 +-- {reme_ai => reme}/core/flow/cmd_flow.py | 0 .../core/flow/expression_flow.py | 11 +- {reme_ai => reme}/core/llm/__init__.py | 6 + {reme_ai => reme}/core/llm/base_llm.py | 181 ++---- {reme_ai => reme}/core/llm/lite_llm.py | 2 - {reme_ai => reme}/core/llm/lite_llm_sync.py | 2 - {reme_ai => reme}/core/llm/openai_llm.py | 2 - {reme_ai => reme}/core/llm/openai_llm_sync.py | 2 - {reme_ai => reme}/core/op/__init__.py | 5 + {reme_ai => reme}/core/op/base_op.py | 227 +++----- {reme_ai => reme}/core/op/base_ray_op.py | 10 +- reme/core/op/base_tool.py | 61 ++ {reme_ai => reme}/core/op/mcp_tool.py | 40 +- {reme_ai => reme}/core/op/parallel_op.py | 4 +- {reme_ai => reme}/core/op/sequential_op.py | 8 +- reme/core/schema/response.py | 4 +- reme/core/schema/service_config.py | 17 +- reme/core/schema/tool_call.py | 5 - {reme_ai => reme}/core/service/__init__.py | 5 + .../core/service/base_service.py | 10 +- {reme_ai => reme}/core/service/cmd_service.py | 10 +- .../core/service/http_service.py | 6 +- {reme_ai => reme}/core/service/mcp_service.py | 14 +- .../core/token_counter/__init__.py | 5 + .../core/token_counter/base_token_counter.py | 3 +- .../core/token_counter/hf_token_counter.py | 2 - .../token_counter/openai_token_counter.py | 4 +- reme/core/utils/__init__.py | 32 ++ {reme_ai => reme}/core/utils/cache_handler.py | 0 .../core/utils/case_converter.py | 0 {reme_ai => reme}/core/utils/common_utils.py | 0 {reme_ai => reme}/core/utils/env_utils.py | 0 .../core/utils/execute_utils.py | 0 {reme_ai => reme}/core/utils/http_client.py | 0 {reme_ai => reme}/core/utils/llm_utils.py | 0 {reme_ai => reme}/core/utils/logger_utils.py | 0 {reme_ai => reme}/core/utils/logo_utils.py | 0 {reme_ai => reme}/core/utils/mcp_client.py | 0 .../core/utils/pydantic_config_parser.py | 0 .../core/utils/pydantic_utils.py | 0 {reme_ai => reme}/core/utils/time.py | 0 .../core/vector_store/__init__.py | 7 + .../core/vector_store/base_vector_store.py | 15 +- .../core/vector_store/chroma_vector_store.py | 8 +- .../core/vector_store/es_vector_store.py | 38 +- .../core/vector_store/local_vector_store.py | 4 +- .../core/vector_store/pgvector_store.py | 26 +- .../core/vector_store/qdrant_vector_store.py | 10 +- reme/reme_app.py | 90 +++ reme_ai/core/__init__.py | 17 - reme_ai/core/application.py | 221 ------- reme_ai/core/context/__init__.py | 16 - reme_ai/core/context/base_context.py | 41 -- reme_ai/core/context/prompt_handler.py | 95 ---- reme_ai/core/context/registry.py | 46 -- reme_ai/core/context/service_context.py | 537 ------------------ reme_ai/core/enumeration/__init__.py | 17 - reme_ai/core/enumeration/chunk_enum.py | 25 - reme_ai/core/enumeration/http_enum.py | 22 - reme_ai/core/enumeration/json_schema_enum.py | 18 - reme_ai/core/enumeration/memory_type.py | 25 - reme_ai/core/enumeration/registry_enum.py | 28 - reme_ai/core/enumeration/role.py | 19 - reme_ai/core/flow/simple_flow.py | 17 - reme_ai/core/main.py | 48 -- reme_ai/core/schema/__init__.py | 42 -- reme_ai/core/schema/memory_node.py | 223 -------- reme_ai/core/schema/message.py | 164 ------ reme_ai/core/schema/request.py | 11 - reme_ai/core/schema/response.py | 11 - reme_ai/core/schema/service_config.py | 113 ---- reme_ai/core/schema/stream_chunk.py | 14 - reme_ai/core/schema/tool_call.py | 228 -------- reme_ai/core/schema/vector_node.py | 15 - reme_ai/core/utils/__init__.py | 47 -- reme_ai/{core => }/reme.py | 1 - 89 files changed, 758 insertions(+), 2523 deletions(-) rename {reme_ai/core => reme}/config/__init__.py (100%) rename {reme_ai/core => reme}/config/default.yaml (88%) rename {reme_ai/core => reme}/config/reme_config_parser.py (76%) rename {reme_ai => reme}/core/context/runtime_context.py (88%) create mode 100644 reme/core/context/service_context.py rename {reme_ai => reme}/core/embedding/__init__.py (65%) rename {reme_ai => reme}/core/embedding/base_embedding_model.py (100%) rename {reme_ai => reme}/core/embedding/openai_embedding_model.py (96%) rename {reme_ai => reme}/core/embedding/openai_embedding_model_sync.py (94%) rename {reme_ai => reme}/core/flow/__init__.py (77%) rename {reme_ai => reme}/core/flow/base_flow.py (80%) rename {reme_ai => reme}/core/flow/cmd_flow.py (100%) rename {reme_ai => reme}/core/flow/expression_flow.py (76%) rename {reme_ai => reme}/core/llm/__init__.py (60%) rename {reme_ai => reme}/core/llm/base_llm.py (67%) rename {reme_ai => reme}/core/llm/lite_llm.py (98%) rename {reme_ai => reme}/core/llm/lite_llm_sync.py (97%) rename {reme_ai => reme}/core/llm/openai_llm.py (98%) rename {reme_ai => reme}/core/llm/openai_llm_sync.py (97%) rename {reme_ai => reme}/core/op/__init__.py (72%) rename {reme_ai => reme}/core/op/base_op.py (61%) rename {reme_ai => reme}/core/op/base_ray_op.py (94%) create mode 100644 reme/core/op/base_tool.py rename {reme_ai => reme}/core/op/mcp_tool.py (75%) rename {reme_ai => reme}/core/op/parallel_op.py (93%) rename {reme_ai => reme}/core/op/sequential_op.py (84%) rename {reme_ai => reme}/core/service/__init__.py (64%) rename {reme_ai => reme}/core/service/base_service.py (76%) rename {reme_ai => reme}/core/service/cmd_service.py (75%) rename {reme_ai => reme}/core/service/http_service.py (95%) rename {reme_ai => reme}/core/service/mcp_service.py (85%) rename {reme_ai => reme}/core/token_counter/__init__.py (58%) rename {reme_ai => reme}/core/token_counter/base_token_counter.py (96%) rename {reme_ai => reme}/core/token_counter/hf_token_counter.py (97%) rename {reme_ai => reme}/core/token_counter/openai_token_counter.py (97%) rename {reme_ai => reme}/core/utils/cache_handler.py (100%) rename {reme_ai => reme}/core/utils/case_converter.py (100%) rename {reme_ai => reme}/core/utils/common_utils.py (100%) rename {reme_ai => reme}/core/utils/env_utils.py (100%) rename reme_ai/core/utils/execute_tuils.py => reme/core/utils/execute_utils.py (100%) rename {reme_ai => reme}/core/utils/http_client.py (100%) rename {reme_ai => reme}/core/utils/llm_utils.py (100%) rename {reme_ai => reme}/core/utils/logger_utils.py (100%) rename {reme_ai => reme}/core/utils/logo_utils.py (100%) rename {reme_ai => reme}/core/utils/mcp_client.py (100%) rename {reme_ai => reme}/core/utils/pydantic_config_parser.py (100%) rename {reme_ai => reme}/core/utils/pydantic_utils.py (100%) rename {reme_ai => reme}/core/utils/time.py (100%) rename {reme_ai => reme}/core/vector_store/__init__.py (62%) rename {reme_ai => reme}/core/vector_store/base_vector_store.py (91%) rename {reme_ai => reme}/core/vector_store/chroma_vector_store.py (99%) rename {reme_ai => reme}/core/vector_store/es_vector_store.py (95%) rename {reme_ai => reme}/core/vector_store/local_vector_store.py (99%) rename {reme_ai => reme}/core/vector_store/pgvector_store.py (96%) rename {reme_ai => reme}/core/vector_store/qdrant_vector_store.py (98%) create mode 100644 reme/reme_app.py delete mode 100644 reme_ai/core/__init__.py delete mode 100644 reme_ai/core/application.py delete mode 100644 reme_ai/core/context/__init__.py delete mode 100644 reme_ai/core/context/base_context.py delete mode 100644 reme_ai/core/context/prompt_handler.py delete mode 100644 reme_ai/core/context/registry.py delete mode 100644 reme_ai/core/context/service_context.py delete mode 100644 reme_ai/core/enumeration/__init__.py delete mode 100644 reme_ai/core/enumeration/chunk_enum.py delete mode 100644 reme_ai/core/enumeration/http_enum.py delete mode 100644 reme_ai/core/enumeration/json_schema_enum.py delete mode 100644 reme_ai/core/enumeration/memory_type.py delete mode 100644 reme_ai/core/enumeration/registry_enum.py delete mode 100644 reme_ai/core/enumeration/role.py delete mode 100644 reme_ai/core/flow/simple_flow.py delete mode 100644 reme_ai/core/main.py delete mode 100644 reme_ai/core/schema/__init__.py delete mode 100644 reme_ai/core/schema/memory_node.py delete mode 100644 reme_ai/core/schema/message.py delete mode 100644 reme_ai/core/schema/request.py delete mode 100644 reme_ai/core/schema/response.py delete mode 100644 reme_ai/core/schema/service_config.py delete mode 100644 reme_ai/core/schema/stream_chunk.py delete mode 100644 reme_ai/core/schema/tool_call.py delete mode 100644 reme_ai/core/schema/vector_node.py delete mode 100644 reme_ai/core/utils/__init__.py rename reme_ai/{core => }/reme.py (99%) diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 3d6fcf96..917a5cd8 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -3,7 +3,7 @@ repos: rev: v6.0.0 hooks: - id: check-ast - exclude: ^(test/|cookbook/|reme_ai/core_old/|reme_ai/mem_agent/|reme_ai/mem_tool/|bench) + exclude: ^(test/|cookbook/|reme_ai/|bench) - id: check-yaml - id: check-xml - id: check-toml @@ -14,18 +14,18 @@ repos: rev: v4.0.0 hooks: - id: add-trailing-comma - exclude: ^(test/|cookbook/|reme_ai/core_old/|reme_ai/mem_agent/|reme_ai/mem_tool/|bench) + exclude: ^(test/|cookbook/|reme_ai/|bench) - repo: https://github.com/psf/black rev: 25.9.0 hooks: - id: black - exclude: ^(test/|cookbook/|reme_ai/core_old/|reme_ai/mem_agent/|reme_ai/mem_tool/|bench) + exclude: ^(test/|cookbook/|reme_ai/|bench) args: [--line-length=120] - repo: https://github.com/PyCQA/flake8 rev: 7.3.0 hooks: - id: flake8 - exclude: ^(test/|cookbook/|reme_ai/core_old/|reme_ai/mem_agent/|reme_ai/mem_tool/|bench) + exclude: ^(test/|cookbook/|reme_ai/|bench) args: [ "--extend-ignore=E203", "--max-line-length=120" @@ -44,9 +44,7 @@ repos: | \.demo$ | \.md$ | \.html$ - | reme_ai/core_old/ - | reme_ai/mem_agent/ - | reme_ai/mem_tool/ + | reme_ai/ | bench ) args: [ @@ -80,6 +78,7 @@ repos: --disable=C3001, --disable=R1702, --disable=R0912, + --max-statements=75, --max-line-length=120, ] - repo: https://github.com/regebro/pyroma diff --git a/reme_ai/core/config/__init__.py b/reme/config/__init__.py similarity index 100% rename from reme_ai/core/config/__init__.py rename to reme/config/__init__.py diff --git a/reme_ai/core/config/default.yaml b/reme/config/default.yaml similarity index 88% rename from reme_ai/core/config/default.yaml rename to reme/config/default.yaml index 00ef4062..836d866f 100644 --- a/reme_ai/core/config/default.yaml +++ b/reme/config/default.yaml @@ -15,15 +15,14 @@ http: llm: default: backend: openai -# model_name: qwen3-30b-a3b-instruct-2507 - model_name: qwen-flash + model_name: qwen3-30b-a3b-instruct-2507 +# model_name: qwen-flash request_interval: 1 temperature: 0.0001 qwen3_max_instruct: backend: openai model_name: qwen3-max -# temperature: 0.6 request_interval: 2 embedding_model: diff --git a/reme_ai/core/config/reme_config_parser.py b/reme/config/reme_config_parser.py similarity index 76% rename from reme_ai/core/config/reme_config_parser.py rename to reme/config/reme_config_parser.py index 798235b2..7a21f806 100644 --- a/reme_ai/core/config/reme_config_parser.py +++ b/reme/config/reme_config_parser.py @@ -1,6 +1,6 @@ """Configuration parser for ReMe framework.""" -from ..utils import PydanticConfigParser +from ..core.utils import PydanticConfigParser class ReMeConfigParser(PydanticConfigParser): diff --git a/reme/core/context/__init__.py b/reme/core/context/__init__.py index 27957fd2..7bbd5869 100644 --- a/reme/core/context/__init__.py +++ b/reme/core/context/__init__.py @@ -3,9 +3,13 @@ from .base_context import BaseContext from .prompt_handler import PromptHandler from .registry_factory import R +from .runtime_context import RuntimeContext +from .service_context import ServiceContext __all__ = [ "BaseContext", "PromptHandler", "R", + "RuntimeContext", + "ServiceContext", ] diff --git a/reme_ai/core/context/runtime_context.py b/reme/core/context/runtime_context.py similarity index 88% rename from reme_ai/core/context/runtime_context.py rename to reme/core/context/runtime_context.py index d7112e1c..44167ca0 100644 --- a/reme_ai/core/context/runtime_context.py +++ b/reme/core/context/runtime_context.py @@ -3,6 +3,7 @@ import asyncio from .base_context import BaseContext +from .service_context import ServiceContext from ..enumeration import ChunkEnum from ..schema import Response, StreamChunk @@ -14,21 +15,23 @@ class RuntimeContext(BaseContext): self, response: Response | None = None, stream_queue: asyncio.Queue | None = None, + service_context: ServiceContext | None = None, **kwargs, ): """Initialize the context with optional response and queue.""" super().__init__(**kwargs) - self.response = response or Response() - self.stream_queue = stream_queue + self.response: Response | None = response or Response() + self.stream_queue: asyncio.Queue | None = stream_queue + self.service_context: ServiceContext | None = service_context @classmethod def from_context(cls, context: "RuntimeContext | None" = None, **kwargs) -> "RuntimeContext": """Create a new context from an existing instance or keywords.""" if context is None: return cls(**kwargs) - - context.update(kwargs) - return context + else: + context.update(kwargs) + return context async def _enqueue(self, chunk: StreamChunk) -> None: """Internal helper to put a chunk into the queue if it exists.""" diff --git a/reme/core/context/service_context.py b/reme/core/context/service_context.py new file mode 100644 index 00000000..5be7b449 --- /dev/null +++ b/reme/core/context/service_context.py @@ -0,0 +1,230 @@ +"""Service context.""" + +import os +from concurrent.futures import ThreadPoolExecutor + +from loguru import logger + +from .base_context import BaseContext +from .registry_factory import R +from ..schema import ServiceConfig +from ..utils import MCPClient, print_logo, PydanticConfigParser, init_logger, load_env, run_coro_safely + + +class ServiceContext(BaseContext): + """Service context.""" + + def __init__( + self, + *args, + llm_api_key: str | None = None, + llm_api_base: str | None = None, + embedding_api_key: str | None = None, + embedding_api_base: str | None = None, + service_config: ServiceConfig | None = None, + parser: type[PydanticConfigParser] | None = None, + config_path: str | None = None, + enable_logo: bool = True, + llm: dict | None = None, + embedding_model: dict | None = None, + vector_store: dict | None = None, + token_counter: dict | None = None, + **kwargs, + ): + super().__init__() + # Set environment variables + load_env() + self._update_env("REME_LLM_API_KEY", llm_api_key) + self._update_env("REME_LLM_BASE_URL", llm_api_base) + self._update_env("REME_EMBEDDING_API_KEY", embedding_api_key) + self._update_env("REME_EMBEDDING_BASE_URL", embedding_api_base) + + # Use default parser if not provided + parser_class = parser if parser is not None else PydanticConfigParser + self.parser = parser_class(ServiceConfig) + + # Service configuration + if service_config is None: + input_args = [] + if config_path: + input_args.append(f"config={config_path}") + if args: + input_args.extend(args) + if kwargs: + input_args.extend([f"{k}={v}" for k, v in kwargs.items()]) + service_config = self.parser.parse_args(*input_args) + self.service_config: ServiceConfig = service_config + + # Initialize logger + if self.service_config.init_logger: + init_logger() + + # Update service config with provided arguments + if llm: + self.update_section_config("llm", **llm) + if embedding_model: + self.update_section_config("embedding_model", **embedding_model) + if token_counter: + self.update_section_config("token_counter", **token_counter) + if vector_store: + self.update_section_config("vector_store", **vector_store) + + # Print the ReMe logo if enabled in configuration. + self.service_config.enable_logo = enable_logo + if self.service_config.enable_logo: + print_logo(service_config=self.service_config) + + # Service configuration and runtime settings + self.language: str = self.service_config.language + self.thread_pool: ThreadPoolExecutor = ThreadPoolExecutor(max_workers=service_config.thread_pool_max_workers) + + # Initialize Ray for distributed computing if configured + if self.service_config.ray_max_workers > 1: + import ray + + ray.init(num_cpus=self.service_config.ray_max_workers) + + from ..llm import BaseLLM + from ..embedding import BaseEmbeddingModel + from ..vector_store import BaseVectorStore + from ..token_counter import BaseTokenCounter + from ..flow import BaseFlow, ExpressionFlow + from ..service import BaseService + + # Initialize LLM instances + self.llms: dict[str, BaseLLM] = {} + for name, config in self.service_config.llm.items(): + self.llms[name] = R.llm[config.backend](model_name=config.model_name, **config.model_extra) + + # Initialize Embedding model instances + self.embedding_models: dict[str, BaseEmbeddingModel] = {} + for name, config in self.service_config.embedding_model.items(): + self.embedding_models[name] = R.embedding_model[config.backend]( + model_name=config.model_name, + **config.model_extra, + ) + + # Initialize Token counter instances + self.token_counters: dict[str, BaseTokenCounter] = {} + for name, config in self.service_config.token_counter.items(): + self.token_counters[name] = R.token_counter[config.backend]( + model_name=config.model_name, + **config.model_extra, + ) + + # Initialize Vector store instances + self.vector_stores: dict[str, BaseVectorStore] = {} + for name, config in self.service_config.vector_store.items(): + self.vector_stores[name] = R.vector_store[config.backend]( + collection_name=config.collection_name, + embedding_model=self.embedding_models[config.embedding_model], + thread_pool=self.thread_pool, + **config.model_extra, + ) + + # Initialize flow instances + self.flows: dict[str, BaseFlow] = {} + for name, flow_cls in R.flow.items(): + if not self._filter_flows(name): + continue + flow: "BaseFlow" = flow_cls(name=name, service_context=self) + self.flows[flow.name] = flow + + # Initialize flow instances from service config + for name, flow_config in self.service_config.flow.items(): + if not self._filter_flows(name): + continue + flow_config.name = name + flow: BaseFlow = ExpressionFlow(flow_config=flow_config, service_context=self) + self.flows[flow.name] = flow + + # Initialize service instance + self.service: BaseService = R.service[self.service_config.backend](service_context=self) + + # MCP server mapping: maps server_name -> {tool_name: ToolCall} + if self.service_config.mcp_servers: + self.mcp_server_mapping: dict[str, dict] = run_coro_safely(self.prepare_mcp_servers()) + else: + self.mcp_server_mapping: dict[str, dict] = {} + + @staticmethod + def _update_env(key: str, value: str | None): + """Update environment variable if value is provided.""" + if value: + os.environ[key] = value + + def update_section_config(self, section_name: str, **kwargs): + """Update a specific section of the service config with new values.""" + section_dict: dict = getattr(self.service_config, section_name) + if "default" not in section_dict: + raise KeyError(f"Default `{section_name}` config not found") + + current_config = section_dict["default"] + section_dict["default"] = current_config.model_copy(update=kwargs, deep=True) + + def _filter_flows(self, name: str) -> bool: + """Filter flows based on enabled_flows and disabled_flows configuration.""" + if self.service_config.enabled_flows: + return name in self.service_config.enabled_flows + elif self.service_config.disabled_flows: + return name not in self.service_config.disabled_flows + else: + return True + + async def prepare_mcp_servers(self): + """Prepare and initialize MCP server connections.""" + mcp_client = MCPClient(config={"mcpServers": self.service_config.mcp_servers}) + for server_name in self.service_config.mcp_servers.keys(): + try: + # Retrieve all available tool calls from this MCP server + tool_calls = await mcp_client.list_tool_calls(server_name=server_name, return_dict=False) + + # Build mapping: tool_name -> ToolCall for quick lookup + self.mcp_server_mapping[server_name] = {tool_call.name: tool_call for tool_call in tool_calls} + + # Log discovered tools for debugging + for tool_call in tool_calls: + logger.info(f"list_tool_calls: {server_name}@{tool_call.name} {tool_call.simple_input_dump()}") + + except Exception as e: + logger.exception(f"list_tool_calls: {server_name} error: {e}") + + async def close(self): + """Close all service components asynchronously.""" + for _, vector_store in self.vector_stores.items(): + await vector_store.close() + + for _, llm in self.llms.items(): + await llm.close() + + for _, embedding_model in self.embedding_models.items(): + await embedding_model.close() + + self.shutdown_thread_pool() + self.shutdown_ray() + + def close_sync(self): + """Close all service components synchronously.""" + for _, vector_store in self.vector_stores.items(): + run_coro_safely(vector_store.close()) + + for _, llm in self.llms.items(): + llm.close_sync() + + for _, embedding_model in self.embedding_models.items(): + embedding_model.close_sync() + + self.shutdown_thread_pool() + self.shutdown_ray() + + def shutdown_thread_pool(self, wait: bool = True): + """Shutdown the thread pool executor.""" + if self.thread_pool: + self.thread_pool.shutdown(wait=wait) + + def shutdown_ray(self, wait: bool = True): + """Shutdown Ray cluster if it was initialized.""" + if self.service_config and self.service_config.ray_max_workers > 1: + import ray + + ray.shutdown(_exiting_interpreter=not wait) diff --git a/reme_ai/core/embedding/__init__.py b/reme/core/embedding/__init__.py similarity index 65% rename from reme_ai/core/embedding/__init__.py rename to reme/core/embedding/__init__.py index c1d92375..f694d065 100644 --- a/reme_ai/core/embedding/__init__.py +++ b/reme/core/embedding/__init__.py @@ -3,9 +3,13 @@ from .base_embedding_model import BaseEmbeddingModel from .openai_embedding_model import OpenAIEmbeddingModel from .openai_embedding_model_sync import OpenAIEmbeddingModelSync +from ..context import R __all__ = [ "BaseEmbeddingModel", "OpenAIEmbeddingModel", "OpenAIEmbeddingModelSync", ] + +R.embedding_model.register("openai")(OpenAIEmbeddingModel) +R.embedding_model.register("openai_sync")(OpenAIEmbeddingModelSync) diff --git a/reme_ai/core/embedding/base_embedding_model.py b/reme/core/embedding/base_embedding_model.py similarity index 100% rename from reme_ai/core/embedding/base_embedding_model.py rename to reme/core/embedding/base_embedding_model.py diff --git a/reme_ai/core/embedding/openai_embedding_model.py b/reme/core/embedding/openai_embedding_model.py similarity index 96% rename from reme_ai/core/embedding/openai_embedding_model.py rename to reme/core/embedding/openai_embedding_model.py index a7e9f0c8..435229b9 100644 --- a/reme_ai/core/embedding/openai_embedding_model.py +++ b/reme/core/embedding/openai_embedding_model.py @@ -6,10 +6,8 @@ from typing import Literal from openai import AsyncOpenAI from .base_embedding_model import BaseEmbeddingModel -from ..context import C -@C.register_embedding_model("openai") class OpenAIEmbeddingModel(BaseEmbeddingModel): """Asynchronous embedding model implementation compatible with OpenAI-style APIs.""" diff --git a/reme_ai/core/embedding/openai_embedding_model_sync.py b/reme/core/embedding/openai_embedding_model_sync.py similarity index 94% rename from reme_ai/core/embedding/openai_embedding_model_sync.py rename to reme/core/embedding/openai_embedding_model_sync.py index 760732cd..cf3aac14 100644 --- a/reme_ai/core/embedding/openai_embedding_model_sync.py +++ b/reme/core/embedding/openai_embedding_model_sync.py @@ -3,10 +3,8 @@ from openai import OpenAI from .openai_embedding_model import OpenAIEmbeddingModel -from ..context import C -@C.register_embedding_model("openai_sync") class OpenAIEmbeddingModelSync(OpenAIEmbeddingModel): """Synchronous embedding model implementation that extends the asynchronous OpenAI model.""" diff --git a/reme_ai/core/flow/__init__.py b/reme/core/flow/__init__.py similarity index 77% rename from reme_ai/core/flow/__init__.py rename to reme/core/flow/__init__.py index e74a2b5c..6d5a053b 100644 --- a/reme_ai/core/flow/__init__.py +++ b/reme/core/flow/__init__.py @@ -3,11 +3,9 @@ from .base_flow import BaseFlow from .cmd_flow import CmdFlow from .expression_flow import ExpressionFlow -from .simple_flow import SimpleFlow __all__ = [ "BaseFlow", "CmdFlow", "ExpressionFlow", - "SimpleFlow", ] diff --git a/reme_ai/core/flow/base_flow.py b/reme/core/flow/base_flow.py similarity index 80% rename from reme_ai/core/flow/base_flow.py rename to reme/core/flow/base_flow.py index 7decce48..b79528f7 100644 --- a/reme_ai/core/flow/base_flow.py +++ b/reme/core/flow/base_flow.py @@ -7,30 +7,25 @@ from abc import ABC, abstractmethod from loguru import logger -from ..context import C, RuntimeContext -from ..enumeration import ChunkEnum, RegistryEnum +from ..context import RuntimeContext, ServiceContext, R +from ..enumeration import ChunkEnum from ..op import BaseOp, SequentialOp, ParallelOp -from ..schema import Response, ToolCall, ToolAttr +from ..schema import Response, ToolCall from ..utils import camel_to_snake, CacheHandler class BaseFlow(ABC): - """Abstract base class for flow execution with caching, streaming, and operation tree management. - - BaseFlow provides a framework for building complex workflows by composing operations - into executable trees. It supports both synchronous and asynchronous execution modes, - response caching, streaming outputs, and automatic tool call schema generation. - """ + """Abstract base class for flow execution with caching, streaming, and operation tree management.""" def __init__( self, name: str = "", - flow_op: BaseOp | None = None, stream: bool = False, raise_exception: bool = True, enable_cache: bool = False, cache_path: str = "cache/flow", cache_expire_hours: float = 0.1, + service_context: ServiceContext | None = None, **kwargs, ): """Initialize flow configuration and execution state.""" @@ -42,11 +37,12 @@ class BaseFlow(ABC): self.enable_cache: bool = enable_cache self.cache_path: str = cache_path self.cache_expire_hours: float = cache_expire_hours + self.service_context: ServiceContext | None = service_context self.flow_params: dict = kwargs - self._flow_op: BaseOp | None = flow_op self._cache: CacheHandler | None = None self._flow_printed: bool = False + self._flow_op: BaseOp | None = None self._tool_call: ToolCall | None = None def _build_tool_call(self) -> ToolCall | None: @@ -82,11 +78,7 @@ class BaseFlow(ABC): return if key := self._compute_cache_key(params): - self.cache.save( - key, - response.model_dump(exclude_none=True), - expire_hours=self.cache_expire_hours, - ) + self.cache.save(key, response.model_dump(exclude_none=True), expire_hours=self.cache_expire_hours) def _print_operation_tree(self, name: str, op: BaseOp, indent: int): """Recursively log the hierarchy of the flow's operation tree.""" @@ -100,19 +92,13 @@ class BaseFlow(ABC): @property def tool_call(self) -> ToolCall | None: """Lazily construct the ToolCall schema describing this flow.""" - if self.flow_op.tool_call: + if hasattr(self.flow_op, "tool_call"): return self.flow_op.tool_call if self._tool_call is None: self._tool_call = self._build_tool_call() if self._tool_call: self._tool_call.name = self._tool_call.name or self.name - self._tool_call.output = self._tool_call.output or { - f"{self.name}_result": ToolAttr( - type="string", - description=f"The execution result of the {self.name}", - ), - } return self._tool_call @property @@ -130,12 +116,6 @@ class BaseFlow(ABC): self._flow_op = self._build_flow() return self._flow_op - @flow_op.setter - def flow_op(self, op: BaseOp): - """Set the root operation of the flow.""" - self._flow_op = op - self._flow_printed = False - @property def async_mode(self) -> bool: """Check if the current flow operation tree is asynchronous.""" @@ -148,11 +128,10 @@ class BaseFlow(ABC): if not lines: raise ValueError("Expression is empty") - env: dict = C.registry_dict[RegistryEnum.OP] if len(lines) > 1: - exec("\n".join(lines[:-1]), {"__builtins__": {}}, env) + exec("\n".join(lines[:-1]), {"__builtins__": {}}, R.op) - result = eval(lines[-1], {"__builtins__": {}}, env) + result = eval(lines[-1], {"__builtins__": {}}, R.op) if not isinstance(result, BaseOp): raise TypeError(f"Expression evaluated to {type(result)}, expected BaseOp") return result @@ -172,30 +151,34 @@ class BaseFlow(ABC): if cached := self._maybe_load_cached(kwargs): return cached - context = RuntimeContext(**kwargs) + context = RuntimeContext(service_context=self.service_context, **kwargs) try: self.print_flow() flow_op: BaseOp = self._build_flow() assert self.flow_op.async_mode, "Async call requires an async flow operation." - await flow_op.call(context=context) - result = context.stream_queue if self.stream else context.response if self.stream: await context.add_stream_done() + return context.stream_queue + + else: + self._maybe_save_cache(kwargs, context.response) + return context.response - self._maybe_save_cache(kwargs, result) - return result except Exception as e: logger.exception(f"[{self.__class__.__name__}] {self.name} async call failed: {e}") if self.raise_exception: raise e + if self.stream: await context.add_stream_chunk_and_type(str(e), ChunkEnum.ERROR) await context.add_stream_done() return context.stream_queue - context.add_response_error(e) - return context.response + + else: + context.add_response_error(e) + return context.response def call_sync(self, **kwargs) -> Response: """Execute the flow synchronously with parameter caching.""" @@ -204,18 +187,20 @@ class BaseFlow(ABC): if cached := self._maybe_load_cached(kwargs): return cached - context = RuntimeContext(**kwargs) + context = RuntimeContext(service_context=self.service_context, **kwargs) try: self.print_flow() flow_op: BaseOp = self._build_flow() assert not self.flow_op.async_mode, "Sync call requires a sync flow operation." - flow_op.call_sync(context=context) + self._maybe_save_cache(kwargs, context.response) return context.response + except Exception as e: logger.exception(f"[{self.__class__.__name__}] {self.name} sync call failed: {e}") if self.raise_exception: raise e + context.add_response_error(e) return context.response diff --git a/reme_ai/core/flow/cmd_flow.py b/reme/core/flow/cmd_flow.py similarity index 100% rename from reme_ai/core/flow/cmd_flow.py rename to reme/core/flow/cmd_flow.py diff --git a/reme_ai/core/flow/expression_flow.py b/reme/core/flow/expression_flow.py similarity index 76% rename from reme_ai/core/flow/expression_flow.py rename to reme/core/flow/expression_flow.py index 5ac8f0d8..b32c9257 100644 --- a/reme_ai/core/flow/expression_flow.py +++ b/reme/core/flow/expression_flow.py @@ -1,6 +1,7 @@ """Expression-based flow implementation driven by configuration objects.""" from .base_flow import BaseFlow +from ..context import ServiceContext from ..op import BaseOp from ..schema import FlowConfig, ToolCall @@ -8,7 +9,7 @@ from ..schema import FlowConfig, ToolCall class ExpressionFlow(BaseFlow): """A flow implementation that constructs operations from a FlowConfig definition.""" - def __init__(self, flow_config: FlowConfig): + def __init__(self, flow_config: FlowConfig, service_context: ServiceContext): """Initialize the flow using settings and metadata from a FlowConfig instance.""" self.flow_config: FlowConfig = flow_config super().__init__( @@ -18,6 +19,7 @@ class ExpressionFlow(BaseFlow): enable_cache=self.flow_config.enable_cache, cache_path=self.flow_config.cache_path, cache_expire_hours=self.flow_config.cache_expire_hours, + service_context=service_context, **flow_config.model_extra, ) @@ -27,4 +29,9 @@ class ExpressionFlow(BaseFlow): def _build_tool_call(self) -> ToolCall: """Construct a tool call representation based on configuration parameters.""" - return ToolCall(**{"description": self.flow_config.description, "parameters": self.flow_config.parameters}) + return ToolCall( + **{ + "description": self.flow_config.description, + "parameters": self.flow_config.parameters, + }, + ) diff --git a/reme_ai/core/llm/__init__.py b/reme/core/llm/__init__.py similarity index 60% rename from reme_ai/core/llm/__init__.py rename to reme/core/llm/__init__.py index 57578f49..1b80641e 100644 --- a/reme_ai/core/llm/__init__.py +++ b/reme/core/llm/__init__.py @@ -5,6 +5,7 @@ from .lite_llm import LiteLLM from .lite_llm_sync import LiteLLMSync from .openai_llm import OpenAILLM from .openai_llm_sync import OpenAILLMSync +from ..context import R __all__ = [ "BaseLLM", @@ -13,3 +14,8 @@ __all__ = [ "OpenAILLM", "OpenAILLMSync", ] + +R.llm.register("litellm")(LiteLLM) +R.llm.register("litellm_sync")(LiteLLMSync) +R.llm.register("openai")(OpenAILLM) +R.llm.register("openai_sync")(OpenAILLMSync) diff --git a/reme_ai/core/llm/base_llm.py b/reme/core/llm/base_llm.py similarity index 67% rename from reme_ai/core/llm/base_llm.py rename to reme/core/llm/base_llm.py index fcd543b6..be7b5e0d 100644 --- a/reme_ai/core/llm/base_llm.py +++ b/reme/core/llm/base_llm.py @@ -1,4 +1,4 @@ -"""Abstract base interface for ReMe LLM implementations.""" +"""Base interface for LLM implementations.""" import asyncio import json @@ -15,16 +15,23 @@ from ..schema import ToolCall class BaseLLM(ABC): - """Abstract base class defining the standard interface for LLM interactions.""" + """Base class for LLM interactions.""" - def __init__(self, model_name: str, max_retries: int = 10, raise_exception: bool = False, request_interval: float = 0.0, **kwargs): - """Initialize the LLM client with model configurations and retry policies. + def __init__( + self, + model_name: str, + max_retries: int = 10, + raise_exception: bool = False, + request_interval: float = 0.0, + **kwargs, + ): + """Initialize LLM client. Args: - model_name: The name of the model to use - max_retries: Maximum number of retry attempts on failure - raise_exception: Whether to raise exceptions or return default values - request_interval: Minimum time interval (in seconds) between consecutive requests. Default is 0.0 (no interval). + model_name: Model name to use + max_retries: Maximum retry attempts on failure + raise_exception: Raise exceptions or return default values + request_interval: Minimum seconds between requests (default: 0.0) **kwargs: Additional model-specific parameters """ self.model_name: str = model_name @@ -33,20 +40,17 @@ class BaseLLM(ABC): self.request_interval: float = request_interval self.kwargs: dict = kwargs - # Request rate control for async operations self._last_request_time: float = 0.0 self._request_lock: asyncio.Lock = asyncio.Lock() @staticmethod def _accumulate_tool_call_chunk(tool_call, ret_tools: list[ToolCall]): - """Assemble incremental tool call fragments into complete ToolCall objects.""" + """Assemble incremental tool call chunks into complete ToolCall objects.""" index = tool_call.index - # Ensure we have a ToolCall object at this index while len(ret_tools) <= index: ret_tools.append(ToolCall(index=index)) - # Accumulate tool call parts (id, name, arguments) if tool_call.id: ret_tools[index].id += tool_call.id @@ -58,7 +62,7 @@ class BaseLLM(ABC): @staticmethod def _validate_and_serialize_tools(ret_tool_calls: list[ToolCall], tools: list[ToolCall]) -> list[dict]: - """Validate tool call integrity and return serialized tool dictionaries.""" + """Validate and serialize tool calls.""" if not ret_tool_calls: return [] @@ -69,10 +73,9 @@ class BaseLLM(ABC): if tool.name not in tool_dict: continue - # First try sanitizing arguments if not tool.sanitize_and_check_argument(): - logger.error(f"Tool call {tool.name} has invalid JSON arguments after sanitization attempt: {tool.arguments}") - raise ValueError(f"Tool call {tool.name} has invalid JSON arguments: {tool.arguments}") + logger.error(f"Invalid JSON arguments in {tool.name}: {tool.arguments}") + raise ValueError(f"Invalid JSON arguments in {tool.name}: {tool.arguments}") validated_tools.append(tool.simple_output_dump()) return validated_tools @@ -86,15 +89,7 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> dict: - """Construct provider-specific parameters for streaming API requests. - - Args: - messages: List of conversation messages - tools: Optional list of tool calls - log_params: Whether to log parameters - model_name: Optional model name to override self.model_name - **kwargs: Additional parameters - """ + """Build provider-specific streaming parameters.""" async def _stream_chat( self, @@ -102,7 +97,7 @@ class BaseLLM(ABC): tools: list[ToolCall] | None, stream_kwargs: dict, ) -> AsyncGenerator[StreamChunk, None]: - """Internal async generator for streaming raw response chunks.""" + """Async generator for streaming response chunks.""" raise NotImplementedError def _stream_chat_sync( @@ -111,7 +106,7 @@ class BaseLLM(ABC): tools: list[ToolCall] | None = None, stream_kwargs: dict | None = None, ) -> Generator[StreamChunk, None, None]: - """Internal synchronous generator for streaming raw response chunks.""" + """Sync generator for streaming response chunks.""" raise NotImplementedError async def stream_chat( @@ -121,22 +116,13 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> AsyncGenerator[StreamChunk, None]: - """Public async interface for streaming chat completions with retries. - - Args: - messages: List of conversation messages - tools: Optional list of tool calls - model_name: Optional model name to override self.model_name - **kwargs: Additional parameters - """ - # Apply request rate limiting if configured + """Stream chat completions with retries.""" if self.request_interval > 0: async with self._request_lock: current_time = time.time() elapsed = current_time - self._last_request_time if elapsed < self.request_interval: - sleep_time = self.request_interval - elapsed - await asyncio.sleep(sleep_time) + await asyncio.sleep(self.request_interval - elapsed) self._last_request_time = time.time() async for chunk in self._stream_chat_impl(messages, tools, model_name, **kwargs): @@ -149,7 +135,7 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> AsyncGenerator[StreamChunk, None]: - """Internal implementation of stream_chat with retry logic.""" + """Stream chat with retry logic.""" stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs) for i in range(self.max_retries): @@ -159,7 +145,7 @@ class BaseLLM(ABC): return except Exception as e: - logger.exception(f"stream chat with model={self.model_name} encounter error with e={e.args}") + logger.exception(f"Stream chat error (model={self.model_name}): {e.args}") if i == self.max_retries - 1: if self.raise_exception: @@ -177,14 +163,7 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> Generator[StreamChunk, None, None]: - """Public synchronous interface for streaming chat completions with retries. - - Args: - messages: List of conversation messages - tools: Optional list of tool calls - model_name: Optional model name to override self.model_name - **kwargs: Additional parameters - """ + """Stream chat completions synchronously with retries.""" stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs) for i in range(self.max_retries): @@ -193,7 +172,7 @@ class BaseLLM(ABC): return except Exception as e: - logger.exception(f"stream chat sync with model={self.model_name} encounter error with e={e.args}") + logger.exception(f"Stream chat sync error (model={self.model_name}): {e.args}") if i == self.max_retries - 1: if self.raise_exception: @@ -212,15 +191,7 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> Message: - """Internal async method to aggregate a full response by consuming the stream. - - Args: - messages: List of conversation messages - tools: Optional list of tool calls - enable_stream_print: Whether to print stream chunks - model_name: Optional model name to override self.model_name - **kwargs: Additional parameters - """ + """Aggregate full response by consuming the stream.""" state = { "enter_think": False, "enter_answer": False, @@ -231,7 +202,6 @@ class BaseLLM(ABC): stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs) async for stream_chunk in self._stream_chat(messages=messages, tools=tools, stream_kwargs=stream_kwargs): - # Process stream chunk if stream_chunk.chunk_type is ChunkEnum.USAGE: if enable_stream_print: print( @@ -280,15 +250,7 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> Message: - """Internal synchronous method to aggregate a full response by consuming the stream. - - Args: - messages: List of conversation messages - tools: Optional list of tool calls - enable_stream_print: Whether to print stream chunks - model_name: Optional model name to override self.model_name - **kwargs: Additional parameters - """ + """Aggregate full response synchronously by consuming the stream.""" state = { "enter_think": False, "enter_answer": False, @@ -299,7 +261,6 @@ class BaseLLM(ABC): stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs) for stream_chunk in self._stream_chat_sync(messages=messages, tools=tools, stream_kwargs=stream_kwargs): - # Process stream chunk if stream_chunk.chunk_type is ChunkEnum.USAGE: if enable_stream_print: print( @@ -350,28 +311,24 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> Message | Any: - """Perform an async chat completion with integrated retries and error handling. - - Args: - messages: List of conversation messages - tools: Optional list of tool calls - enable_stream_print: Whether to print stream chunks - callback_fn: Optional callback function to process the result - default_value: Default value to return on error - model_name: Optional model name to override self.model_name - **kwargs: Additional parameters - """ - # Apply request rate limiting if configured + """Chat completion with retries and error handling.""" if self.request_interval > 0: async with self._request_lock: current_time = time.time() elapsed = current_time - self._last_request_time if elapsed < self.request_interval: - sleep_time = self.request_interval - elapsed - await asyncio.sleep(sleep_time) + await asyncio.sleep(self.request_interval - elapsed) self._last_request_time = time.time() - return await self._chat_impl(messages, tools, enable_stream_print, callback_fn, default_value, model_name, **kwargs) + return await self._chat_impl( + messages, + tools, + enable_stream_print, + callback_fn, + default_value, + model_name, + **kwargs, + ) async def _chat_impl( self, @@ -383,8 +340,7 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> Message | Any: - """Internal implementation of chat with retry and error handling logic.""" - # Use the provided model_name or fall back to self.model_name + """Chat with retry and error handling logic.""" effective_model = model_name if model_name is not None else self.model_name for i in range(self.max_retries): @@ -399,19 +355,16 @@ class BaseLLM(ABC): return callback_fn(result) if callback_fn else result except Exception as e: - # Check if this is an inappropriate content error error_message = str(e.args[0]) if e.args else str(e) is_inappropriate_content = "inappropriate content" in error_message.lower() is_rate_limit_error = ( - "request rate increased too quickly" in error_message.lower() or - "exceeded your current quota" in error_message.lower() or - "insufficient_quota" in error_message.lower() + "request rate increased too quickly" in error_message.lower() + or "exceeded your current quota" in error_message.lower() + or "insufficient_quota" in error_message.lower() ) if is_inappropriate_content: - logger.error(f"chat with model={effective_model} detected inappropriate content error") - logger.error("=" * 80) - logger.error("Full message content that triggered the error:") + logger.error(f"Inappropriate content detected (model={effective_model})") logger.error("=" * 80) for idx, msg in enumerate(messages): logger.error(f"Message {idx + 1} [role={msg.role}]:") @@ -422,15 +375,16 @@ class BaseLLM(ABC): logger.error(f"Tool calls: {msg.tool_calls}") logger.error("-" * 80) logger.error("=" * 80) - # Return empty Message immediately without retrying return Message(role=Role.ASSISTANT, content="") if is_rate_limit_error: - logger.warning(f"chat with model={effective_model} hit rate limit, sleeping for 60s before retry (attempt {i + 1}/{self.max_retries})") + logger.warning( + f"Rate limit hit (model={effective_model}), sleeping 60s (attempt {i + 1}/{self.max_retries})", + ) await asyncio.sleep(60) continue - logger.exception(f"chat with model={effective_model} encounter error with e={e.args}") + logger.exception(f"Chat error (model={effective_model}): {e.args}") if i == self.max_retries - 1: if self.raise_exception: @@ -450,18 +404,7 @@ class BaseLLM(ABC): model_name: str | None = None, **kwargs, ) -> Message | Any: - """Perform a synchronous chat completion with integrated retries and error handling. - - Args: - messages: List of conversation messages - tools: Optional list of tool calls - enable_stream_print: Whether to print stream chunks - callback_fn: Optional callback function to process the result - default_value: Default value to return on error - model_name: Optional model name to override self.model_name - **kwargs: Additional parameters - """ - # Use the provided model_name or fall back to self.model_name + """Chat completion synchronously with retries and error handling.""" effective_model = model_name if model_name is not None else self.model_name for i in range(self.max_retries): @@ -476,19 +419,16 @@ class BaseLLM(ABC): return callback_fn(result) if callback_fn else result except Exception as e: - # Check if this is an inappropriate content error error_message = str(e.args[0]) if e.args else str(e) is_inappropriate_content = "inappropriate content" in error_message.lower() is_rate_limit_error = ( - "request rate increased too quickly" in error_message.lower() or - "exceeded your current quota" in error_message.lower() or - "insufficient_quota" in error_message.lower() + "request rate increased too quickly" in error_message.lower() + or "exceeded your current quota" in error_message.lower() + or "insufficient_quota" in error_message.lower() ) if is_inappropriate_content: - logger.error(f"chat sync with model={effective_model} detected inappropriate content error") - logger.error("=" * 80) - logger.error("Full message content that triggered the error:") + logger.error(f"Inappropriate content detected (model={effective_model})") logger.error("=" * 80) for idx, msg in enumerate(messages): logger.error(f"Message {idx + 1} [role={msg.role}]:") @@ -499,15 +439,16 @@ class BaseLLM(ABC): logger.error(f"Tool calls: {msg.tool_calls}") logger.error("-" * 80) logger.error("=" * 80) - # Return empty Message immediately without retrying return Message(role=Role.ASSISTANT, content="") if is_rate_limit_error: - logger.warning(f"chat sync with model={effective_model} hit rate limit, sleeping for 60s before retry (attempt {i + 1}/{self.max_retries})") + logger.warning( + f"Rate limit hit (model={effective_model}), sleeping 60s (attempt {i + 1}/{self.max_retries})", + ) time.sleep(60) continue - logger.exception(f"chat sync with model={effective_model} encounter error with e={e.args}") + logger.exception(f"Chat sync error (model={effective_model}): {e.args}") if i == self.max_retries - 1: if self.raise_exception: @@ -518,7 +459,7 @@ class BaseLLM(ABC): return default_value async def close(self): - """Release any asynchronous resources or connections held by the client.""" + """Release async resources.""" def close_sync(self): - """Release any synchronous resources or connections held by the client.""" + """Release sync resources.""" diff --git a/reme_ai/core/llm/lite_llm.py b/reme/core/llm/lite_llm.py similarity index 98% rename from reme_ai/core/llm/lite_llm.py rename to reme/core/llm/lite_llm.py index 458fe1c0..5663702b 100644 --- a/reme_ai/core/llm/lite_llm.py +++ b/reme/core/llm/lite_llm.py @@ -7,14 +7,12 @@ import litellm from loguru import logger from .base_llm import BaseLLM -from ..context import C from ..enumeration import ChunkEnum from ..schema import Message from ..schema import StreamChunk from ..schema import ToolCall -@C.register_llm("litellm") class LiteLLM(BaseLLM): """Async LLM implementation using LiteLLM to support multiple providers.""" diff --git a/reme_ai/core/llm/lite_llm_sync.py b/reme/core/llm/lite_llm_sync.py similarity index 97% rename from reme_ai/core/llm/lite_llm_sync.py rename to reme/core/llm/lite_llm_sync.py index 11fbc7af..778eaed0 100644 --- a/reme_ai/core/llm/lite_llm_sync.py +++ b/reme/core/llm/lite_llm_sync.py @@ -5,14 +5,12 @@ from typing import Generator import litellm from .lite_llm import LiteLLM -from ..context import C from ..enumeration import ChunkEnum from ..schema import Message from ..schema import StreamChunk from ..schema import ToolCall -@C.register_llm("litellm_sync") class LiteLLMSync(LiteLLM): """Synchronous LiteLLM client for executing chat completions and streaming responses.""" diff --git a/reme_ai/core/llm/openai_llm.py b/reme/core/llm/openai_llm.py similarity index 98% rename from reme_ai/core/llm/openai_llm.py rename to reme/core/llm/openai_llm.py index f15dcb8e..ea647e07 100644 --- a/reme_ai/core/llm/openai_llm.py +++ b/reme/core/llm/openai_llm.py @@ -7,14 +7,12 @@ from loguru import logger from openai import AsyncOpenAI from .base_llm import BaseLLM -from ..context import C from ..enumeration import ChunkEnum from ..schema import Message from ..schema import StreamChunk from ..schema import ToolCall -@C.register_llm("openai") class OpenAILLM(BaseLLM): """Asynchronous LLM client for OpenAI-compatible APIs supporting streaming completions and tool execution.""" diff --git a/reme_ai/core/llm/openai_llm_sync.py b/reme/core/llm/openai_llm_sync.py similarity index 97% rename from reme_ai/core/llm/openai_llm_sync.py rename to reme/core/llm/openai_llm_sync.py index a2bcfee3..86dd119d 100644 --- a/reme_ai/core/llm/openai_llm_sync.py +++ b/reme/core/llm/openai_llm_sync.py @@ -5,14 +5,12 @@ from typing import Generator from openai import OpenAI from .openai_llm import OpenAILLM -from ..context import C from ..enumeration import ChunkEnum from ..schema import Message from ..schema import StreamChunk from ..schema import ToolCall -@C.register_llm("openai_sync") class OpenAILLMSync(OpenAILLM): """Synchronous LLM client for OpenAI-compatible APIs, inheriting from OpenAILLM.""" diff --git a/reme_ai/core/op/__init__.py b/reme/core/op/__init__.py similarity index 72% rename from reme_ai/core/op/__init__.py rename to reme/core/op/__init__.py index a48bf276..53a6f379 100644 --- a/reme_ai/core/op/__init__.py +++ b/reme/core/op/__init__.py @@ -2,14 +2,19 @@ from .base_op import BaseOp from .base_ray_op import BaseRayOp +from .base_tool import BaseTool from .mcp_tool import MCPTool from .parallel_op import ParallelOp from .sequential_op import SequentialOp +from ..context import R __all__ = [ "BaseOp", "BaseRayOp", + "BaseTool", "MCPTool", "ParallelOp", "SequentialOp", ] + +R.op.register("mcp_tool")(MCPTool) diff --git a/reme_ai/core/op/base_op.py b/reme/core/op/base_op.py similarity index 61% rename from reme_ai/core/op/base_op.py rename to reme/core/op/base_op.py index 357db44e..ef469d0f 100644 --- a/reme_ai/core/op/base_op.py +++ b/reme/core/op/base_op.py @@ -3,22 +3,23 @@ import asyncio import copy import inspect +from abc import ABCMeta from pathlib import Path -from typing import Callable, Optional +from typing import Callable, Optional, Any from loguru import logger from tqdm import tqdm -from ..context import RuntimeContext, PromptHandler, C +from ..context import RuntimeContext, PromptHandler, ServiceContext from ..embedding import BaseEmbeddingModel from ..llm import BaseLLM -from ..schema import ToolCall, ToolAttr, Response +from ..schema import Response from ..token_counter import BaseTokenCounter from ..utils import camel_to_snake, CacheHandler, timer from ..vector_store import BaseVectorStore -class BaseOp: +class BaseOp(metaclass=ABCMeta): """Base operator class for LLM workflow execution and composition.""" def __new__(cls, *args, **kwargs): @@ -34,6 +35,7 @@ class BaseOp: async_mode: bool = True, language: str = "", prompt_name: str = "", + prompt_path: str = "", llm: str | BaseLLM = "default", embedding_model: str | BaseEmbeddingModel = "default", vector_store: str | BaseVectorStore = "default", @@ -44,7 +46,6 @@ class BaseOp: sub_ops: dict[str, "BaseOp"] | list["BaseOp"] | Optional["BaseOp"] = None, input_mapping: dict[str, str] | None = None, output_mapping: dict[str, str] | None = None, - save_response_result: bool = False, enable_sync_thread_pool: bool = True, max_retries: int = 1, raise_exception: bool = False, @@ -53,8 +54,8 @@ class BaseOp: """Initialize operator configurations and internal state.""" self.name = name or camel_to_snake(self.__class__.__name__) self.async_mode = async_mode - self.language = language or C.language - self.prompt = self._get_prompt_handler(prompt_name) + self.language = language + self.prompt = self._get_prompt_handler(prompt_name, prompt_path) self._llm = llm self._embedding_model = embedding_model @@ -64,12 +65,12 @@ class BaseOp: self.enable_cache = enable_cache self.cache_path = cache_path self.cache_expire_hours = cache_expire_hours - self.sub_ops: list[BaseOp] = [] + + self.sub_ops: list["BaseOp"] = [] self.add_sub_ops(sub_ops) self.input_mapping = input_mapping self.output_mapping = output_mapping - self.save_response_result = save_response_result self.enable_sync_thread_pool = enable_sync_thread_pool self.max_retries = max(1, max_retries) self.raise_exception = raise_exception @@ -78,86 +79,29 @@ class BaseOp: self._pending_tasks: list = [] self.context: RuntimeContext | None = None self._cache: CacheHandler | None = None - self._tool_call: ToolCall | None = None - def _get_prompt_handler(self, prompt_name: str) -> PromptHandler: + def _get_prompt_handler(self, prompt_name: str, prompt_path: str) -> PromptHandler: """Load prompt configuration from the associated YAML file.""" - path = Path(inspect.getfile(self.__class__)) - path = path.with_stem(prompt_name) if prompt_name else path + if prompt_path: + path = Path(prompt_path) + else: + path = Path(inspect.getfile(self.__class__)) + if prompt_name: + path = path.with_stem(prompt_name) return PromptHandler(language=self.language).load_prompt_by_file(path.with_suffix(".yaml")) - def _build_tool_call(self) -> ToolCall | None: - """Build and return the tool call schema; override in subclasses.""" - - def _validate_inputs(self): - """Ensure all required tool inputs are present in context.""" - if self.tool_call is not None: - parameters = self.tool_call.parameters - if parameters.type == "object" and parameters.properties: - required_list = parameters.required or [] - required_keys = {k: (k in required_list) for k in parameters.properties.keys()} - self.context.validate_required_keys(required_keys, self.name) - - def _handle_failure(self, e: Exception, attempt: int): + def _handle_failure(self, e: Exception, attempt: int) -> str | None: """Log failures and handle final retry logic.""" message = f"[{self.__class__.__name__}] {self.name} failed (attempt {attempt + 1}): {e}" if attempt == self.max_retries - 1: logger.exception(message) if self.raise_exception: raise e - - if self.tool_call is not None: - self.output = f"{self.name} failed: {e}" + return f"{self.name} failed: {e}" else: logger.warning(message) - - @property - def tool_call(self) -> ToolCall | None: - """Lazily construct and return the tool call metadata.""" - if self._tool_call is None: - self._tool_call = self._build_tool_call() - if self._tool_call is None: - return None - - self._tool_call.name = self._tool_call.name or self.name - if not self._tool_call.output.properties: - self._tool_call.output.properties = { - f"{self.name}_result": ToolAttr(type="string", description=f"Execution result of {self.name}"), - } - return self._tool_call - - @property - def input_dict(self) -> dict: - """Extract required and optional inputs from context based on schema.""" - parameters = self.tool_call.parameters - if parameters.type != "object" or not parameters.properties: - return {} - required_keys = set(parameters.required or []) - return {k: self.context[k] for k in parameters.properties.keys() if (k in required_keys or k in self.context)} - - @property - def output(self): - """Get the single output value from context.""" - output_properties = self.tool_call.output.properties - if not output_properties: return None - keys = list(output_properties.keys()) - if len(keys) >= 1 and keys[0] in self.context: - return self.context[keys[0]] - else: - return None - - @output.setter - def output(self, value): - """Set the single output value into context.""" - output_properties = self.tool_call.output.properties - if not output_properties: - return - - keys = list(output_properties.keys()) - self.context[keys[0]] = value - @property def cache(self) -> CacheHandler: """Access the operator-specific cache handler.""" @@ -166,119 +110,120 @@ class BaseOp: self._cache = CacheHandler(f"{self.cache_path}/{self.name}") return self._cache + @property + def service_context(self) -> ServiceContext: + """Access the service context.""" + return self.context.service_context + @property def llm(self) -> BaseLLM: """Get the LLM instance from ServiceContext.""" if isinstance(self._llm, str): - self._llm = C.get_llm(self._llm) + self._llm = self.service_context.llms[self._llm] return self._llm @property def embedding_model(self) -> BaseEmbeddingModel: """Get the embedding model instance from ServiceContext.""" if isinstance(self._embedding_model, str): - self._embedding_model = C.get_embedding_model(self._embedding_model) + self._embedding_model = self.service_context.embedding_models[self._embedding_model] return self._embedding_model @property def vector_store(self) -> BaseVectorStore: """Lazily initialize and return the vector store instance.""" if isinstance(self._vector_store, str): - self._vector_store = C.get_vector_store(self._vector_store) + self._vector_store = self.service_context.vector_stores[self._vector_store] return self._vector_store @property def token_counter(self) -> BaseTokenCounter: """Get the token counter instance from ServiceContext.""" if isinstance(self._token_counter, str): - self._token_counter = C.get_token_counter(self._token_counter) + self._token_counter = self.service_context.token_counters[self._token_counter] return self._token_counter @property def service_metadata(self) -> dict: """Get service configuration metadata.""" - return C.service_config.model_extra + return self.service_context.service_config.model_extra @property def response(self) -> Response: - """Get the response object.""" + """Access the response object.""" return self.context.response - def set_tool_call(self, tool_call: ToolCall | dict): - """Set the tool call.""" - if isinstance(tool_call, dict): - self._tool_call = ToolCall(**tool_call) - elif isinstance(tool_call, ToolCall): - self._tool_call = tool_call - else: - raise ValueError(f"Invalid tool call: {tool_call}") - - self._tool_call.name = self._tool_call.name or self.name - if not self._tool_call.output.properties: - self._tool_call.output.properties = { - f"{self.name}_result": ToolAttr(type="string", description=f"Execution result of {self.name}"), - } - - def set_language(self, language: str): - """Set the language.""" - self.language = language - return self - def before_execute_sync(self): """Prepare context and validate before sync execution.""" self.context.apply_mapping(self.input_mapping) - self._validate_inputs() + + async def before_execute(self): + """Prepare context and validate before async execution.""" + self.context.apply_mapping(self.input_mapping) def execute_sync(self): """Define core sync logic in subclasses.""" - def after_execute_sync(self): - """Finalize context and mappings after sync execution.""" - self.context.apply_mapping(self.output_mapping) - if self.tool_call is not None and self.save_response_result: - self.context.response.answer = self.output - - async def before_execute(self): - """Prepare context and validate before async execution.""" - self.before_execute_sync() - async def execute(self): """Define core async logic in subclasses.""" - async def after_execute(self): + def after_execute_sync(self, response: Any): + """Finalize context and mappings after sync execution.""" + self.context.apply_mapping(self.output_mapping) + if response is not None: + if isinstance(response, dict): + for k, v in response.items(): + if k == "answer": + self.response.answer = v + elif k == "success": + self.response.success = v.lower() == "true" + else: + self.response.metadata[k] = v + else: + self.response.answer = response + return response + + async def after_execute(self, output: Any): """Finalize context and mappings after async execution.""" - self.after_execute_sync() + return self.after_execute_sync(output) @timer def call_sync(self, context: RuntimeContext = None, **kwargs): """Execute the operator synchronously with retry logic.""" self.context = RuntimeContext.from_context(context, **kwargs) + response = None for i in range(self.max_retries): try: self.before_execute_sync() - self.execute_sync() - self.after_execute_sync() + response = self.execute_sync() + response = self.after_execute_sync(response) break except Exception as e: - self._handle_failure(e, i) - return self.output if self.tool_call is not None else None + response = self._handle_failure(e, i) + return response + + @timer async def call(self, context: RuntimeContext = None, **kwargs): """Execute the operator asynchronously with retry logic.""" self.context = RuntimeContext.from_context(context, **kwargs) + response = None for i in range(self.max_retries): try: await self.before_execute() - await self.execute() - await self.after_execute() + response = await self.execute() + response = await self.after_execute(response) break except Exception as e: - self._handle_failure(e, i) - return self.output if self.tool_call is not None else None + response = self._handle_failure(e, i) + return response def submit_sync_task(self, fn: Callable, *args, **kwargs) -> "BaseOp": """Submit a task to the thread pool or local queue.""" - task = C.thread_pool.submit(fn, *args, **kwargs) if self.enable_sync_thread_pool else (fn, args, kwargs) + if self.enable_sync_thread_pool: + task = self.service_context.thread_pool.submit(fn, *args, **kwargs) + else: + task = (fn, args, kwargs) self._pending_tasks.append(task) return self @@ -292,26 +237,32 @@ class BaseOp: """Wait for all pending sync tasks and return flattened results.""" results = [] for task in tqdm(self._pending_tasks, desc=task_desc or self.name): - res = task.result() if self.enable_sync_thread_pool else task[0](*task[1], **task[2]) - if res: - results.extend(res if isinstance(res, list) else [res]) + if self.enable_sync_thread_pool: + result = task.result() + else: + result = task[0](*task[1], **task[2]) + if result: + if isinstance(result, list): + results.extend(result) + else: + results.append(result) self._pending_tasks.clear() return results async def join_async_tasks(self, return_exceptions: bool = True) -> list: """Wait for all pending async tasks and aggregate results.""" - try: - raw_results = await asyncio.gather(*self._pending_tasks, return_exceptions=return_exceptions) - results = [] - for res in raw_results: - if isinstance(res, Exception): - logger.error(f"[{self.__class__.__name__}] Async task failed: {res}") - continue - if res: - results.extend(res if isinstance(res, list) else [res]) - return results - finally: - self._pending_tasks.clear() + raw_results = await asyncio.gather(*self._pending_tasks, return_exceptions=return_exceptions) + results = [] + for result in raw_results: + if isinstance(result, Exception): + logger.error(f"[{self.__class__.__name__}] Async task failed: {result}") + elif result: + if isinstance(result, list): + results.extend(result) + else: + result.append(result) + self._pending_tasks.clear() + return results def add_sub_ops(self, sub_ops: dict[str, "BaseOp"] | list["BaseOp"] | Optional["BaseOp"]): """Add child operators to this operator's sub_ops.""" diff --git a/reme_ai/core/op/base_ray_op.py b/reme/core/op/base_ray_op.py similarity index 94% rename from reme_ai/core/op/base_ray_op.py rename to reme/core/op/base_ray_op.py index 7a8fe113..84c0ed25 100644 --- a/reme_ai/core/op/base_ray_op.py +++ b/reme/core/op/base_ray_op.py @@ -8,14 +8,14 @@ from loguru import logger from tqdm import tqdm from .base_op import BaseOp -from ..context import BaseContext, C +from ..context import BaseContext _RAY_IMPORT_ERROR = None try: import ray -except ImportError as e: - _RAY_IMPORT_ERROR = e +except ImportError as _e: + _RAY_IMPORT_ERROR = _e ray = None @@ -35,7 +35,7 @@ class BaseRayOp(BaseOp, metaclass=ABCMeta): def submit_and_join_ray_task(self, fn: Callable, parallel_key: str = "", task_desc: str = "", **kwargs) -> list: """Divide data into chunks and execute them across Ray workers.""" - max_workers = C.service_config.ray_max_workers + max_workers = self.service_context.ray_max_workers self._ray_task_list.clear() # Automatically detect the key containing the list to parallelize @@ -94,7 +94,7 @@ class BaseRayOp(BaseOp, metaclass=ABCMeta): def submit_ray_task(self, fn, *args, **kwargs): """Submit a single Ray task to the task list for later execution.""" if not ray.is_initialized(): - ray.init(num_cpus=C.service_config.ray_max_workers, ignore_reinit_error=True) + ray.init(num_cpus=self.service_context.ray_max_workers, ignore_reinit_error=True) remote_fn = ray.remote(fn) task = remote_fn.remote(*args, **kwargs) diff --git a/reme/core/op/base_tool.py b/reme/core/op/base_tool.py new file mode 100644 index 00000000..9daf033a --- /dev/null +++ b/reme/core/op/base_tool.py @@ -0,0 +1,61 @@ +"""Base class for tools""" + +from abc import ABCMeta + +from . import BaseOp +from ..schema import ToolCall + + +class BaseTool(BaseOp, metaclass=ABCMeta): + """Base class for tools""" + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self._tool_call: ToolCall | None = None + + def _build_tool_call(self) -> ToolCall: + """Build and return the tool call schema; override in subclasses.""" + + def _validate_inputs(self): + """Validate the inputs.""" + parameters = self.tool_call.parameters + if parameters.type == "object" and parameters.properties: + required_list = parameters.required or [] + required_keys = {k: (k in required_list) for k in parameters.properties.keys()} + self.context.validate_required_keys(required_keys, self.name) + + @property + def tool_call(self) -> ToolCall | None: + """Get the tool call schema.""" + if self._tool_call is None: + self._tool_call = self._build_tool_call() + if self._tool_call is None: + return None + + self._tool_call.name = self._tool_call.name or self.name + return self._tool_call + + def set_tool_call(self, tool_call: ToolCall | dict): + """Set the tool call schema.""" + if isinstance(tool_call, dict): + self._tool_call = ToolCall(**tool_call) + elif isinstance(tool_call, ToolCall): + self._tool_call = tool_call + else: + raise ValueError(f"Invalid tool call: {tool_call}") + + self._tool_call.name = self._tool_call.name or self.name + + @property + def input_dict(self) -> dict: + """Get the input dict.""" + parameters = self.tool_call.parameters + if parameters.type != "object" or not parameters.properties: + return {} + required_keys = set(parameters.required or []) + return {k: self.context[k] for k in parameters.properties.keys() if (k in required_keys or k in self.context)} + + def before_execute_sync(self): + """Hook before execute""" + super().before_execute_sync() + self._validate_inputs() diff --git a/reme_ai/core/op/mcp_tool.py b/reme/core/op/mcp_tool.py similarity index 75% rename from reme_ai/core/op/mcp_tool.py rename to reme/core/op/mcp_tool.py index 57cbcbd4..c59e105c 100644 --- a/reme_ai/core/op/mcp_tool.py +++ b/reme/core/op/mcp_tool.py @@ -2,25 +2,20 @@ from typing import List -from .base_op import BaseOp -from ..context import C +from mcp.types import CallToolResult, TextContent + +from .base_tool import BaseTool from ..schema import ToolCall from ..utils import MCPClient -@C.register_op() -class MCPTool(BaseOp): - """Operator for calling remote MCP (Model Context Protocol) tools. - - This class enables integration with external MCP servers to execute tools - and retrieve their results. It supports parameter customization and retry logic. - """ +class MCPTool(BaseTool): + """Operator for calling remote MCP (Model Context Protocol) tools.""" def __init__( self, mcp_server: str = "", tool_name: str = "", - save_response_result: bool = True, parameter_required: List[str] | None = None, parameter_optional: List[str] | None = None, parameter_deleted: List[str] | None = None, @@ -29,13 +24,7 @@ class MCPTool(BaseOp): raise_exception: bool = False, **kwargs, ): - - super().__init__( - save_response_result=save_response_result, - max_retries=max_retries, - raise_exception=raise_exception, - **kwargs, - ) + super().__init__(max_retries=max_retries, raise_exception=raise_exception, **kwargs) self.mcp_server: str = mcp_server self.tool_name: str = tool_name @@ -43,12 +32,12 @@ class MCPTool(BaseOp): self.parameter_optional: List[str] | None = parameter_optional self.parameter_deleted: List[str] | None = parameter_deleted self.timeout: float | None = timeout - # Example MCP marketplace: https://bailian.console.aliyun.com/?tab=mcp#/mcp-market - self._client = MCPClient(C.service_config.mcp_servers) + # Example MCP marketplace: https://bailian.console.aliyun.com/?tab=mcp#/mcp-market + self._client = MCPClient(self.service_context.service_config.mcp_servers) def _build_tool_call(self) -> ToolCall: - tool_call_dict = C.mcp_server_mapping[self.mcp_server] + tool_call_dict = self.service_context.mcp_server_mapping[self.mcp_server] tool_call: ToolCall = tool_call_dict[self.tool_name].model_copy(deep=True) # Initialize required list if not exists @@ -74,9 +63,16 @@ class MCPTool(BaseOp): return tool_call async def execute(self): - self.output = await self._client.call_tool( + tool_result: CallToolResult = await self._client.call_tool( server_name=self.mcp_server, tool_name=self.tool_name, arguments=self.input_dict, - parse_text_result=True, ) + self.context.tool_result = tool_result + + text_result = [] + for block in tool_result.content: + if isinstance(block, TextContent): + text_result.append(block.text) + output: str = "\n".join(text_result) + return output diff --git a/reme_ai/core/op/parallel_op.py b/reme/core/op/parallel_op.py similarity index 93% rename from reme_ai/core/op/parallel_op.py rename to reme/core/op/parallel_op.py index 18b84d0c..8bca5790 100644 --- a/reme_ai/core/op/parallel_op.py +++ b/reme/core/op/parallel_op.py @@ -11,14 +11,14 @@ class ParallelOp(BaseOp): for op in self.sub_ops: assert op.async_mode self.submit_async_task(op.call, context=self.context) - await self.join_async_tasks() + return await self.join_async_tasks() def execute_sync(self): """Executes all sub-operations concurrently using synchronous task management.""" for op in self.sub_ops: assert not op.async_mode self.submit_sync_task(op.call_sync, context=self.context) - self.join_sync_tasks() + return self.join_sync_tasks() def __lshift__(self, op: dict[str, BaseOp] | list[BaseOp] | BaseOp): """Raises RuntimeError as the shift operator is not supported for parallel operations.""" diff --git a/reme_ai/core/op/sequential_op.py b/reme/core/op/sequential_op.py similarity index 84% rename from reme_ai/core/op/sequential_op.py rename to reme/core/op/sequential_op.py index 3dabb0c9..fed44dc3 100644 --- a/reme_ai/core/op/sequential_op.py +++ b/reme/core/op/sequential_op.py @@ -8,15 +8,19 @@ class SequentialOp(BaseOp): async def execute(self): """Executes sub-operations sequentially using asynchronous awaits.""" + result = None for op in self.sub_ops: assert op.async_mode - await op.call(context=self.context) + result = await op.call(context=self.context) + return result def execute_sync(self): """Executes sub-operations sequentially in a synchronous blocking manner.""" + result = None for op in self.sub_ops: assert not op.async_mode - op.call_sync(context=self.context) + result = op.call_sync(context=self.context) + return result def __lshift__(self, op: dict[str, BaseOp] | list[BaseOp] | BaseOp): """Raises RuntimeError as the left shift operator is not supported.""" diff --git a/reme/core/schema/response.py b/reme/core/schema/response.py index 3104bc6e..fe753232 100644 --- a/reme/core/schema/response.py +++ b/reme/core/schema/response.py @@ -1,11 +1,13 @@ """Defines the standardized data structure for model output responses.""" +from typing import Any + from pydantic import Field, BaseModel class Response(BaseModel): """Represents a structured response containing the execution result, status, and metadata.""" - answer: str | dict | list = Field(default="") + answer: str | Any = Field(default="") success: bool = Field(default=True) metadata: dict = Field(default_factory=dict) diff --git a/reme/core/schema/service_config.py b/reme/core/schema/service_config.py index e1ae9df7..95461196 100644 --- a/reme/core/schema/service_config.py +++ b/reme/core/schema/service_config.py @@ -1,7 +1,6 @@ """Configuration schemas for service components using Pydantic models.""" import os -from typing import Dict, List from pydantic import BaseModel, Field, ConfigDict @@ -99,15 +98,15 @@ class ServiceConfig(BaseModel): thread_pool_max_workers: int = Field(default=16) ray_max_workers: int = Field(default=-1) init_logger: bool = Field(default=True) - disabled_flows: List[str] = Field(default_factory=list) - enabled_flows: List[str] = Field(default_factory=list) - mcp_servers: Dict[str, dict] = Field(default_factory=dict) + disabled_flows: list[str] = Field(default_factory=list) + enabled_flows: list[str] = Field(default_factory=list) + mcp_servers: dict[str, dict] = Field(default_factory=dict) mcp: MCPConfig = Field(default_factory=MCPConfig) http: HttpConfig = Field(default_factory=HttpConfig) cmd: CmdConfig = Field(default_factory=CmdConfig) - flow: Dict[str, FlowConfig] = Field(default_factory=dict) - llm: Dict[str, LLMConfig] = Field(default_factory=dict) - embedding_model: Dict[str, EmbeddingModelConfig] = Field(default_factory=dict) - vector_store: Dict[str, VectorStoreConfig] = Field(default_factory=dict) - token_counter: Dict[str, TokenCounterConfig] = Field(default_factory=dict) + flow: dict[str, FlowConfig] = Field(default_factory=dict) + llm: dict[str, LLMConfig] = Field(default_factory=dict) + embedding_model: dict[str, EmbeddingModelConfig] = Field(default_factory=dict) + vector_store: dict[str, VectorStoreConfig] = Field(default_factory=dict) + token_counter: dict[str, TokenCounterConfig] = Field(default_factory=dict) diff --git a/reme/core/schema/tool_call.py b/reme/core/schema/tool_call.py index 70355d60..e7ba9b78 100644 --- a/reme/core/schema/tool_call.py +++ b/reme/core/schema/tool_call.py @@ -103,11 +103,6 @@ class ToolCall(BaseModel): description="Specification for input parameters", ) - output: ToolAttr = Field( - default_factory=lambda: ToolAttr(type="object", properties={}), - description="Specification for the execution result (Schema)", - ) - @model_validator(mode="before") @classmethod def init_tool_call(cls, data: dict) -> dict: diff --git a/reme_ai/core/service/__init__.py b/reme/core/service/__init__.py similarity index 64% rename from reme_ai/core/service/__init__.py rename to reme/core/service/__init__.py index e9f00a65..1cd4ae58 100644 --- a/reme_ai/core/service/__init__.py +++ b/reme/core/service/__init__.py @@ -4,6 +4,7 @@ from .base_service import BaseService from .cmd_service import CmdService from .http_service import HttpService from .mcp_service import MCPService +from ..context import R __all__ = [ "BaseService", @@ -11,3 +12,7 @@ __all__ = [ "HttpService", "MCPService", ] + +R.service.register("cmd")(CmdService) +R.service.register("http")(HttpService) +R.service.register("mcp")(MCPService) diff --git a/reme_ai/core/service/base_service.py b/reme/core/service/base_service.py similarity index 76% rename from reme_ai/core/service/base_service.py rename to reme/core/service/base_service.py index 28c82fb3..9496391c 100644 --- a/reme_ai/core/service/base_service.py +++ b/reme/core/service/base_service.py @@ -5,7 +5,7 @@ from abc import ABC, abstractmethod from loguru import logger from pydantic import BaseModel -from ..context import C +from ..context import ServiceContext from ..flow import BaseFlow from ..schema import ToolCall from ..utils import create_pydantic_model @@ -14,8 +14,10 @@ from ..utils import create_pydantic_model class BaseService(ABC): """Abstract base class for services that integrate and execute flows.""" - def __init__(self, **kwargs): + def __init__(self, service_context: ServiceContext, **kwargs): """Initialize the base service.""" + self.service_context: ServiceContext = service_context + self.service_config = self.service_context.service_config self.kwargs = kwargs @abstractmethod @@ -32,10 +34,10 @@ class BaseService(ABC): def run(self): """Initialize and integrate all flows registered in the global context.""" flow_names: list[str] = [] - for _, flow in C.flow_dict.items(): + for flow in self.service_context.flows.values(): flow_name = self.integrate_flow(flow) if flow_name: flow_names.append(flow_name) if flow_names: - logger.info(f"integrate {','.join(flow_names)}") + logger.info(f"Integrated {','.join(flow_names)}") diff --git a/reme_ai/core/service/cmd_service.py b/reme/core/service/cmd_service.py similarity index 75% rename from reme_ai/core/service/cmd_service.py rename to reme/core/service/cmd_service.py index c02efa20..18d147a6 100644 --- a/reme_ai/core/service/cmd_service.py +++ b/reme/core/service/cmd_service.py @@ -3,12 +3,10 @@ from loguru import logger from .base_service import BaseService -from ..context import C from ..flow import CmdFlow, BaseFlow from ..utils.common_utils import run_coro_safely -@C.register_service("cmd") class CmdService(BaseService): """Service implementation for handling command flow execution logic.""" @@ -19,16 +17,16 @@ class CmdService(BaseService): def integrate_flow(self, flow: BaseFlow) -> str | None: """Integrate the workflow configuration into the command service.""" - self._cmd_flow = CmdFlow(flow=C.service_config.flow) + self._cmd_flow = CmdFlow(flow=self.service_config.cmd.flow) def run(self): """Execute the command flow in either asynchronous or synchronous mode.""" super().run() - + kwargs = self.service_config.cmd.model_extra if self._cmd_flow.async_mode: - response = run_coro_safely(self._cmd_flow.call(**C.service_config.cmd.model_extra)) + response = run_coro_safely(self._cmd_flow.call(**kwargs)) else: - response = self._cmd_flow.call_sync(**C.service_config.cmd.model_extra) + response = self._cmd_flow.call_sync(**kwargs) if response.answer: logger.info(f"response.answer={response.answer}") diff --git a/reme_ai/core/service/http_service.py b/reme/core/service/http_service.py similarity index 95% rename from reme_ai/core/service/http_service.py rename to reme/core/service/http_service.py index 8ef5d4ce..694dd3bd 100644 --- a/reme_ai/core/service/http_service.py +++ b/reme/core/service/http_service.py @@ -9,20 +9,18 @@ from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import StreamingResponse from .base_service import BaseService -from ..context import C from ..flow import BaseFlow from ..schema import Response from ..utils.common_utils import execute_stream_task -@C.register_service("http") class HttpService(BaseService): """Expose flows via HTTP REST and SSE endpoints.""" def __init__(self, **kwargs): """Initialize FastAPI app with CORS and health checks.""" super().__init__(**kwargs) - self.app = FastAPI(title=C.service_config.app_name) + self.app = FastAPI(title=self.service_config.app_name) self.app.add_middleware( CORSMiddleware, allow_origins=["*"], @@ -75,7 +73,7 @@ class HttpService(BaseService): def run(self): """Start the Uvicorn server.""" super().run() - cfg = C.service_config.http + cfg = self.service_config.http uvicorn.run( self.app, host=cfg.host, diff --git a/reme_ai/core/service/mcp_service.py b/reme/core/service/mcp_service.py similarity index 85% rename from reme_ai/core/service/mcp_service.py rename to reme/core/service/mcp_service.py index 7ac62a26..65183230 100644 --- a/reme_ai/core/service/mcp_service.py +++ b/reme/core/service/mcp_service.py @@ -1,23 +1,19 @@ """Model Context Protocol (MCP) service implementation.""" -from typing import Any - from fastmcp import FastMCP from fastmcp.tools import FunctionTool from .base_service import BaseService -from ..context import C from ..flow import BaseFlow -@C.register_service("mcp") class MCPService(BaseService): """Expose flows as Model Context Protocol (MCP) tools.""" - def __init__(self, **kwargs: Any): + def __init__(self, **kwargs): """Initialize FastMCP instance with service settings.""" super().__init__(**kwargs) - self.mcp = FastMCP(name=C.service_config.app_name) + self.mcp = FastMCP(name=self.service_config.app_name) def integrate_flow(self, flow: BaseFlow) -> str | None: """Register a non-streaming flow as an MCP tool.""" @@ -45,12 +41,8 @@ class MCPService(BaseService): def run(self): """Run the MCP server with specified transport protocol.""" super().run() - cfg = C.service_config.mcp - + cfg = self.service_config.mcp run_args: dict = {"transport": cfg.transport, "show_banner": False, **cfg.model_extra} - - # Add network settings for non-stdio transports if cfg.transport != "stdio": run_args.update({"host": cfg.host, "port": cfg.port}) - self.mcp.run(**run_args) diff --git a/reme_ai/core/token_counter/__init__.py b/reme/core/token_counter/__init__.py similarity index 58% rename from reme_ai/core/token_counter/__init__.py rename to reme/core/token_counter/__init__.py index a9b50826..f0cb5e28 100644 --- a/reme_ai/core/token_counter/__init__.py +++ b/reme/core/token_counter/__init__.py @@ -3,9 +3,14 @@ from .base_token_counter import BaseTokenCounter from .hf_token_counter import HFTokenCounter from .openai_token_counter import OpenAITokenCounter +from ..context import R __all__ = [ "BaseTokenCounter", "HFTokenCounter", "OpenAITokenCounter", ] + +R.token_counter.register("base")(BaseTokenCounter) +R.token_counter.register("hf")(HFTokenCounter) +R.token_counter.register("openai")(OpenAITokenCounter) diff --git a/reme_ai/core/token_counter/base_token_counter.py b/reme/core/token_counter/base_token_counter.py similarity index 96% rename from reme_ai/core/token_counter/base_token_counter.py rename to reme/core/token_counter/base_token_counter.py index 76ccdb78..6f98cd02 100644 --- a/reme_ai/core/token_counter/base_token_counter.py +++ b/reme/core/token_counter/base_token_counter.py @@ -2,13 +2,12 @@ import math import re + from loguru import logger -from ..context import C from ..schema import Message, ToolCall -@C.register_token_counter("base") class BaseTokenCounter: """A rule-based token counter for Chinese and non-Chinese text.""" diff --git a/reme_ai/core/token_counter/hf_token_counter.py b/reme/core/token_counter/hf_token_counter.py similarity index 97% rename from reme_ai/core/token_counter/hf_token_counter.py rename to reme/core/token_counter/hf_token_counter.py index 0bad4c78..4dd072a9 100644 --- a/reme_ai/core/token_counter/hf_token_counter.py +++ b/reme/core/token_counter/hf_token_counter.py @@ -5,11 +5,9 @@ import os from loguru import logger from .base_token_counter import BaseTokenCounter -from ..context import C from ..schema import Message, ToolCall -@C.register_token_counter("hf") class HFTokenCounter(BaseTokenCounter): """Token counter using transformers.AutoTokenizer.apply_chat_template.""" diff --git a/reme_ai/core/token_counter/openai_token_counter.py b/reme/core/token_counter/openai_token_counter.py similarity index 97% rename from reme_ai/core/token_counter/openai_token_counter.py rename to reme/core/token_counter/openai_token_counter.py index 5793e88a..672f8a45 100644 --- a/reme_ai/core/token_counter/openai_token_counter.py +++ b/reme/core/token_counter/openai_token_counter.py @@ -1,13 +1,13 @@ """Token counting implementation for OpenAI-compatible models.""" import json + from loguru import logger + from .base_token_counter import BaseTokenCounter -from ..context import C from ..schema import Message, ToolCall -@C.register_token_counter("openai") class OpenAITokenCounter(BaseTokenCounter): """Token counter for OpenAI models using tiktoken.""" diff --git a/reme/core/utils/__init__.py b/reme/core/utils/__init__.py index ce9d9b2d..3d54e4e4 100644 --- a/reme/core/utils/__init__.py +++ b/reme/core/utils/__init__.py @@ -1,7 +1,39 @@ """utils""" +from .cache_handler import CacheHandler +from .case_converter import snake_to_camel, camel_to_snake +from .common_utils import run_coro_safely, execute_stream_task +from .env_utils import load_env +from .execute_utils import exec_code, run_shell_command +from .http_client import HttpClient +from .llm_utils import extract_content, format_messages, deduplicate_memories +from .logger_utils import init_logger +from .logo_utils import print_logo +from .mcp_client import MCPClient +from .pydantic_config_parser import PydanticConfigParser +from .pydantic_utils import create_pydantic_model from .singleton import singleton +from .time import timer, get_now_time __all__ = [ + "CacheHandler", + "snake_to_camel", + "camel_to_snake", + "run_coro_safely", + "execute_stream_task", + "load_env", + "exec_code", + "run_shell_command", + "HttpClient", + "extract_content", + "format_messages", + "deduplicate_memories", + "init_logger", + "print_logo", + "MCPClient", + "PydanticConfigParser", + "create_pydantic_model", "singleton", + "timer", + "get_now_time", ] diff --git a/reme_ai/core/utils/cache_handler.py b/reme/core/utils/cache_handler.py similarity index 100% rename from reme_ai/core/utils/cache_handler.py rename to reme/core/utils/cache_handler.py diff --git a/reme_ai/core/utils/case_converter.py b/reme/core/utils/case_converter.py similarity index 100% rename from reme_ai/core/utils/case_converter.py rename to reme/core/utils/case_converter.py diff --git a/reme_ai/core/utils/common_utils.py b/reme/core/utils/common_utils.py similarity index 100% rename from reme_ai/core/utils/common_utils.py rename to reme/core/utils/common_utils.py diff --git a/reme_ai/core/utils/env_utils.py b/reme/core/utils/env_utils.py similarity index 100% rename from reme_ai/core/utils/env_utils.py rename to reme/core/utils/env_utils.py diff --git a/reme_ai/core/utils/execute_tuils.py b/reme/core/utils/execute_utils.py similarity index 100% rename from reme_ai/core/utils/execute_tuils.py rename to reme/core/utils/execute_utils.py diff --git a/reme_ai/core/utils/http_client.py b/reme/core/utils/http_client.py similarity index 100% rename from reme_ai/core/utils/http_client.py rename to reme/core/utils/http_client.py diff --git a/reme_ai/core/utils/llm_utils.py b/reme/core/utils/llm_utils.py similarity index 100% rename from reme_ai/core/utils/llm_utils.py rename to reme/core/utils/llm_utils.py diff --git a/reme_ai/core/utils/logger_utils.py b/reme/core/utils/logger_utils.py similarity index 100% rename from reme_ai/core/utils/logger_utils.py rename to reme/core/utils/logger_utils.py diff --git a/reme_ai/core/utils/logo_utils.py b/reme/core/utils/logo_utils.py similarity index 100% rename from reme_ai/core/utils/logo_utils.py rename to reme/core/utils/logo_utils.py diff --git a/reme_ai/core/utils/mcp_client.py b/reme/core/utils/mcp_client.py similarity index 100% rename from reme_ai/core/utils/mcp_client.py rename to reme/core/utils/mcp_client.py diff --git a/reme_ai/core/utils/pydantic_config_parser.py b/reme/core/utils/pydantic_config_parser.py similarity index 100% rename from reme_ai/core/utils/pydantic_config_parser.py rename to reme/core/utils/pydantic_config_parser.py diff --git a/reme_ai/core/utils/pydantic_utils.py b/reme/core/utils/pydantic_utils.py similarity index 100% rename from reme_ai/core/utils/pydantic_utils.py rename to reme/core/utils/pydantic_utils.py diff --git a/reme_ai/core/utils/time.py b/reme/core/utils/time.py similarity index 100% rename from reme_ai/core/utils/time.py rename to reme/core/utils/time.py diff --git a/reme_ai/core/vector_store/__init__.py b/reme/core/vector_store/__init__.py similarity index 62% rename from reme_ai/core/vector_store/__init__.py rename to reme/core/vector_store/__init__.py index 79500294..bd8a1869 100644 --- a/reme_ai/core/vector_store/__init__.py +++ b/reme/core/vector_store/__init__.py @@ -6,6 +6,7 @@ from .es_vector_store import ESVectorStore from .local_vector_store import LocalVectorStore from .pgvector_store import PGVectorStore from .qdrant_vector_store import QdrantVectorStore +from ..context import R __all__ = [ "BaseVectorStore", @@ -15,3 +16,9 @@ __all__ = [ "PGVectorStore", "QdrantVectorStore", ] + +R.vector_store.register("chroma")(ChromaVectorStore) +R.vector_store.register("es")(ESVectorStore) +R.vector_store.register("local")(LocalVectorStore) +R.vector_store.register("pgvector")(PGVectorStore) +R.vector_store.register("qdrant")(QdrantVectorStore) diff --git a/reme_ai/core/vector_store/base_vector_store.py b/reme/core/vector_store/base_vector_store.py similarity index 91% rename from reme_ai/core/vector_store/base_vector_store.py rename to reme/core/vector_store/base_vector_store.py index a4a8ca8e..40a84e99 100644 --- a/reme_ai/core/vector_store/base_vector_store.py +++ b/reme/core/vector_store/base_vector_store.py @@ -3,11 +3,11 @@ import asyncio from abc import ABC, abstractmethod from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor from functools import partial -from reme_ai.core.context import C -from reme_ai.core.embedding import BaseEmbeddingModel -from reme_ai.core.schema import VectorNode +from ..embedding import BaseEmbeddingModel +from ..schema import VectorNode class BaseVectorStore(ABC): @@ -17,20 +17,19 @@ class BaseVectorStore(ABC): self, collection_name: str, embedding_model: BaseEmbeddingModel, + thread_pool: ThreadPoolExecutor, **kwargs, ): """Initialize the vector store with a collection name and an embedding model.""" - if embedding_model is None: - raise ValueError("embedding_model is required") self.collection_name: str = collection_name self.embedding_model: BaseEmbeddingModel = embedding_model + self.thread_pool: ThreadPoolExecutor = thread_pool self.kwargs: dict = kwargs - @staticmethod - async def _run_sync_in_executor(sync_func: Callable, *args, **kwargs): + async def _run_sync_in_executor(self, sync_func: Callable, *args, **kwargs): """Run a synchronous function in the context-defined thread pool executor.""" loop = asyncio.get_running_loop() - return await loop.run_in_executor(C.thread_pool, partial(sync_func, *args, **kwargs)) + return await loop.run_in_executor(self.thread_pool, partial(sync_func, *args, **kwargs)) # noqa async def get_node_embedding(self, node: VectorNode) -> VectorNode: """Generate and assign embedding for a single vector node.""" diff --git a/reme_ai/core/vector_store/chroma_vector_store.py b/reme/core/vector_store/chroma_vector_store.py similarity index 99% rename from reme_ai/core/vector_store/chroma_vector_store.py rename to reme/core/vector_store/chroma_vector_store.py index 567ac3b1..eaddd0e7 100644 --- a/reme_ai/core/vector_store/chroma_vector_store.py +++ b/reme/core/vector_store/chroma_vector_store.py @@ -5,7 +5,6 @@ from typing import Any from loguru import logger from .base_vector_store import BaseVectorStore -from ..context import C from ..embedding import BaseEmbeddingModel from ..schema import VectorNode @@ -20,7 +19,6 @@ except ImportError as e: Settings = None -@C.register_vector_store("chroma") class ChromaVectorStore(BaseVectorStore): """ChromaDB-based vector store implementation for local or remote storage.""" @@ -142,7 +140,7 @@ class ChromaVectorStore(BaseVectorStore): # ChromaDB requires separate conditions combined with $and return [ {k: {"$gte": v[0]}}, - {k: {"$lte": v[1]}} + {k: {"$lte": v[1]}}, ] if isinstance(v, dict): chroma_condition = {} @@ -239,8 +237,8 @@ class ChromaVectorStore(BaseVectorStore): try: self.client.delete_collection(name=collection_name) return True - except Exception as e: - logger.warning(f"Failed to delete collection {collection_name}: {e}") + except Exception as _e: + logger.warning(f"Failed to delete collection {collection_name}: {_e}") return False deleted = await self._run_sync_in_executor(_delete) diff --git a/reme_ai/core/vector_store/es_vector_store.py b/reme/core/vector_store/es_vector_store.py similarity index 95% rename from reme_ai/core/vector_store/es_vector_store.py rename to reme/core/vector_store/es_vector_store.py index 39989277..49054afc 100644 --- a/reme_ai/core/vector_store/es_vector_store.py +++ b/reme/core/vector_store/es_vector_store.py @@ -9,7 +9,6 @@ from typing import Any from loguru import logger from .base_vector_store import BaseVectorStore -from ..context import C from ..embedding import BaseEmbeddingModel from ..schema import VectorNode @@ -24,7 +23,6 @@ except ImportError as e: async_bulk = None -@C.register_vector_store("es") class ESVectorStore(BaseVectorStore): """Elasticsearch-based vector store for dense vector storage and kNN search.""" @@ -265,14 +263,16 @@ class ESVectorStore(BaseVectorStore): # New syntax: [start, end] represents a range query if isinstance(value, list) and len(value) == 2: # Range query: field >= value[0] AND field <= value[1] - filter_conditions.append({ - "range": { - f"metadata.{key}": { - "gte": value[0], - "lte": value[1] - } - } - }) + filter_conditions.append( + { + "range": { + f"metadata.{key}": { + "gte": value[0], + "lte": value[1], + }, + }, + }, + ) else: # Exact match filter_conditions.append({"term": {f"metadata.{key}": value}}) @@ -461,14 +461,16 @@ class ESVectorStore(BaseVectorStore): # New syntax: [start, end] represents a range query if isinstance(value, list) and len(value) == 2: # Range query: field >= value[0] AND field <= value[1] - filter_conditions.append({ - "range": { - f"metadata.{key}": { - "gte": value[0], - "lte": value[1] - } - } - }) + filter_conditions.append( + { + "range": { + f"metadata.{key}": { + "gte": value[0], + "lte": value[1], + }, + }, + }, + ) else: # Exact match filter_conditions.append({"term": {f"metadata.{key}": value}}) diff --git a/reme_ai/core/vector_store/local_vector_store.py b/reme/core/vector_store/local_vector_store.py similarity index 99% rename from reme_ai/core/vector_store/local_vector_store.py rename to reme/core/vector_store/local_vector_store.py index 7533fa89..24e1977c 100644 --- a/reme_ai/core/vector_store/local_vector_store.py +++ b/reme/core/vector_store/local_vector_store.py @@ -6,12 +6,10 @@ from pathlib import Path from loguru import logger from .base_vector_store import BaseVectorStore -from ..context import C from ..embedding import BaseEmbeddingModel from ..schema import VectorNode -@C.register_vector_store("local") class LocalVectorStore(BaseVectorStore): """Local file system-based vector store using JSON files and manual cosine similarity.""" @@ -110,7 +108,7 @@ class LocalVectorStore(BaseVectorStore): return False try: # Try numeric comparison - if not (value[0] <= node_value <= value[1]): + if not value[0] <= node_value <= value[1]: return False except TypeError: # If comparison fails, the filter doesn't match diff --git a/reme_ai/core/vector_store/pgvector_store.py b/reme/core/vector_store/pgvector_store.py similarity index 96% rename from reme_ai/core/vector_store/pgvector_store.py rename to reme/core/vector_store/pgvector_store.py index 882ced3e..7d1b3401 100644 --- a/reme_ai/core/vector_store/pgvector_store.py +++ b/reme/core/vector_store/pgvector_store.py @@ -7,7 +7,6 @@ from typing import Any from loguru import logger from .base_vector_store import BaseVectorStore -from ..context import C from ..embedding import BaseEmbeddingModel from ..schema import VectorNode @@ -22,7 +21,6 @@ except ImportError as e: Pool = None -@C.register_vector_store("pgvector") class PGVectorStore(BaseVectorStore): """Vector store implementation using PostgreSQL and pgvector for efficient similarity search.""" @@ -39,10 +37,10 @@ class PGVectorStore(BaseVectorStore): raise ValueError("Table name cannot be empty") if len(name) > 63: raise ValueError(f"Table name too long: {len(name)} characters (max 63)") - if not re.match(r'^[a-zA-Z_][a-zA-Z0-9_]*$', name): + if not re.match(r"^[a-zA-Z_][a-zA-Z0-9_]*$", name): raise ValueError( f"Invalid table name: {name}. Must start with letter or underscore, " - "and contain only alphanumeric characters and underscores." + "and contain only alphanumeric characters and underscores.", ) def __init__( @@ -295,8 +293,10 @@ class PGVectorStore(BaseVectorStore): for key, value in filters.items(): # Sanitize key to prevent SQL injection (only allow alphanumeric and underscore) - if not key.replace('_', '').replace('.', '').isalnum(): - raise ValueError(f"Invalid metadata key: {key}. Only alphanumeric characters, underscore and dot are allowed.") + if not key.replace("_", "").replace(".", "").isalnum(): + raise ValueError( + f"Invalid metadata key: {key}. Only alphanumeric characters, underscore and dot are allowed.", + ) # New syntax: [start, end] represents a range query if isinstance(value, list) and len(value) == 2: @@ -305,13 +305,12 @@ class PGVectorStore(BaseVectorStore): if isinstance(value[0], (int, float)) and isinstance(value[1], (int, float)): # Numeric range query conditions.append( - f"(metadata->>'{key}')::numeric >= ${param_idx} AND (metadata->>'{key}')::numeric <= ${param_idx + 1}" + f"(metadata->>'{key}')::numeric >= ${param_idx} AND " + f"(metadata->>'{key}')::numeric <= ${param_idx + 1}", ) else: # Text range query (works for strings, timestamps, etc.) - conditions.append( - f"metadata->>'{key}' >= ${param_idx} AND metadata->>'{key}' <= ${param_idx + 1}" - ) + conditions.append(f"metadata->>'{key}' >= ${param_idx} AND metadata->>'{key}' <= ${param_idx + 1}") params.extend([value[0], value[1]]) param_idx += 2 else: @@ -341,12 +340,9 @@ class PGVectorStore(BaseVectorStore): # Adjust parameter indices in filter clause to account for $1 being used by vector_str if filter_clause: - # Replace from highest index to lowest to avoid conflicts for i in range(len(filter_params), 0, -1): - old_placeholder = f"${i}" new_placeholder = f"${i + 1}" - # Use word boundary to ensure we only replace exact matches (e.g., $1 not $10) - filter_clause = re.sub(rf'\${i}\b', new_placeholder, filter_clause) + filter_clause = re.sub(rf"\${i}\b", new_placeholder, filter_clause) async with pool.acquire() as conn: sql = f""" @@ -417,7 +413,7 @@ class PGVectorStore(BaseVectorStore): async with pool.acquire() as conn: result = await conn.execute(f"DELETE FROM {self.collection_name}") - logger.info(f"Deleted all documents from {self.collection_name}") + logger.info(f"Deleted all documents from {self.collection_name} result={result}") async def update(self, nodes: VectorNode | list[VectorNode], **kwargs): """Update existing vector nodes with new content, embeddings, or metadata.""" diff --git a/reme_ai/core/vector_store/qdrant_vector_store.py b/reme/core/vector_store/qdrant_vector_store.py similarity index 98% rename from reme_ai/core/vector_store/qdrant_vector_store.py rename to reme/core/vector_store/qdrant_vector_store.py index d0b4a8aa..5ca74cd0 100644 --- a/reme_ai/core/vector_store/qdrant_vector_store.py +++ b/reme/core/vector_store/qdrant_vector_store.py @@ -5,7 +5,6 @@ from typing import Any from loguru import logger from .base_vector_store import BaseVectorStore -from ..context import C from ..embedding import BaseEmbeddingModel from ..schema import VectorNode @@ -36,7 +35,6 @@ except ImportError as e: VectorParams = None -@C.register_vector_store("qdrant") class QdrantVectorStore(BaseVectorStore): """Vector store implementation using Qdrant for dense vector search.""" @@ -274,7 +272,7 @@ class QdrantVectorStore(BaseVectorStore): logger.warning( f"Qdrant does not support range queries for non-numeric values. " f"Skipping range filter for key '{key}' with values {value}. " - f"Consider using numeric timestamps instead." + f"Consider using numeric timestamps instead.", ) elif isinstance(value, dict) and ("gte" in value or "lte" in value): range_params = {} @@ -284,7 +282,8 @@ class QdrantVectorStore(BaseVectorStore): range_params["gte"] = value["gte"] else: logger.warning( - f"Qdrant range filter for key '{key}' requires numeric gte value, got {type(value['gte']).__name__}. Skipping." + f"Qdrant range filter for key '{key}' requires numeric gte value, " + f"got {type(value['gte']).__name__}. Skipping.", ) continue if "lte" in value: @@ -292,7 +291,8 @@ class QdrantVectorStore(BaseVectorStore): range_params["lte"] = value["lte"] else: logger.warning( - f"Qdrant range filter for key '{key}' requires numeric lte value, got {type(value['lte']).__name__}. Skipping." + f"Qdrant range filter for key '{key}' requires numeric lte value, " + f"got {type(value['lte']).__name__}. Skipping.", ) continue diff --git a/reme/reme_app.py b/reme/reme_app.py new file mode 100644 index 00000000..483cd5bb --- /dev/null +++ b/reme/reme_app.py @@ -0,0 +1,90 @@ +"""ReMe application classes for simplified configuration and execution.""" + +import asyncio +import sys + +from reme.core.utils import execute_stream_task +from .config import ReMeConfigParser +from .core.context import ServiceContext +from .core.flow import BaseFlow +from .core.schema import Response + + +class ReMeApp: + """ReMe application with config file support and flow execution methods.""" + + def __init__( + self, + *args, + llm_api_key: str | None = None, + llm_api_base: str | None = None, + embedding_api_key: str | None = None, + embedding_api_base: str | None = None, + enable_logo: bool = True, + **kwargs, + ): + self.service_context = ServiceContext( + *args, + llm_api_key=llm_api_key, + llm_api_base=llm_api_base, + embedding_api_key=embedding_api_key, + embedding_api_base=embedding_api_base, + service_config=None, + parser=ReMeConfigParser, + config_path=None, + enable_logo=enable_logo, + **kwargs, + ) + + async def __aenter__(self): + """Async context manager entry.""" + return self + + def __enter__(self): + """Context manager entry.""" + return self + + async def __aexit__(self, exc_type=None, exc_val=None, exc_tb=None): + """Async context manager exit.""" + await self.service_context.close() + return False + + def __exit__(self, exc_type=None, exc_val=None, exc_tb=None): + """Context manager exit.""" + self.service_context.close_sync() + return False + + async def execute_flow(self, name: str, **kwargs) -> Response: + """Execute a flow with the given name and parameters.""" + assert name in self.service_context.flows, f"Flow {name} not found" + flow: BaseFlow = self.service_context.flows[name] + return await flow.call(**kwargs) + + async def execute_stream_flow(self, name: str, **kwargs): + """Execute a stream flow with the given name and parameters.""" + assert name in self.service_context.flows, f"Flow {name} not found" + flow: BaseFlow = self.service_context.flows[name] + assert flow.stream is True, "non-stream flow is not supported in execute_stream_flow!" + stream_queue = asyncio.Queue() + task = asyncio.create_task(flow.call(stream_queue=stream_queue, **kwargs)) + async for chunk in execute_stream_task( + stream_queue=stream_queue, + task=task, + task_name=name, + as_bytes=False, + ): + yield chunk + + def run_service(self): + """Run the configured service (HTTP, MCP, or CMD).""" + self.service_context.service.run() + + +def main(): + """Main entry point for running ReMe application from command line.""" + with ReMeApp(*sys.argv[1:]) as app: + app.run_service() + + +if __name__ == "__main__": + main() diff --git a/reme_ai/core/__init__.py b/reme_ai/core/__init__.py deleted file mode 100644 index 8eab5792..00000000 --- a/reme_ai/core/__init__.py +++ /dev/null @@ -1,17 +0,0 @@ -"""Core module for ReMe AI framework.""" - -# pylint: disable=wrong-import-position -# flake8: noqa: F401 - -from . import config -from . import context -from . import embedding -from . import enumeration -from . import flow -from . import llm -from . import op -from . import schema -from . import service -from . import token_counter -from . import utils -from . import vector_store diff --git a/reme_ai/core/application.py b/reme_ai/core/application.py deleted file mode 100644 index 29eb846e..00000000 --- a/reme_ai/core/application.py +++ /dev/null @@ -1,221 +0,0 @@ -"""Main application module for managing ReMe AI service lifecycle and flow execution.""" - -import asyncio -import os - -from .context import C -from .flow import BaseFlow -from .schema import ServiceConfig, Response -from .utils import PydanticConfigParser, init_logger, execute_stream_task, run_coro_safely, load_env - - -class Application: - """ - Main application class for managing the lifecycle of ReMe AI services. - - Handles initialization, configuration, service management, and flow execution - for both synchronous and asynchronous contexts. - """ - - def __init__( - self, - *args, - llm_api_key: str | None = None, - llm_api_base: str | None = None, - embedding_api_key: str | None = None, - embedding_api_base: str | None = None, - service_config: ServiceConfig | None = None, - parser: type[PydanticConfigParser] | None = None, - config_path: str | None = None, - enable_logo: bool = True, - llm: dict | None = None, - embedding_model: dict | None = None, - vector_store: dict | None = None, - token_counter: dict | None = None, - **kwargs, - ): - """ - Initialize the Application with configuration settings. - - Args: - *args: Additional arguments passed to parser. Examples: - - "llm.default.model_name=qwen3-30b-a3b-thinking-2507" - - "llm.default.backend=openai_compatible" - - "llm.default.temperature=0.6" - - "embedding_model.default.model_name=text-embedding-v4" - - "embedding_model.default.backend=openai_compatible" - - "embedding_model.default.dimensions=1024" - - "vector_store.default.backend=memory" - - "vector_store.default.embedding_model=default" - llm_api_key: API key for LLM service - llm_api_base: Base URL for LLM service - embedding_api_key: API key for embedding service - embedding_api_base: Base URL for embedding service - service_config: Pre-built service configuration object - parser: Custom parser class for configuration (defaults to PydanticConfigParser) - config_path: Path to configuration file - enable_logo: Whether to display the ReMe logo on startup - llm: LLM configuration dictionary - embedding_model: Embedding model configuration dictionary - vector_store: Vector store configuration dictionary - token_counter: Token counter configuration dictionary - **kwargs: Additional keyword arguments passed to parser. Same format as args but as kwargs. Examples: - - **{"llm.default.model_name": "qwen3-30b-a3b-thinking-2507"} - """ - - load_env() - self._update_env("REME_LLM_API_KEY", llm_api_key) - self._update_env("REME_LLM_BASE_URL", llm_api_base) - self._update_env("REME_EMBEDDING_API_KEY", embedding_api_key) - self._update_env("REME_EMBEDDING_BASE_URL", embedding_api_base) - - # Use default parser if not provided - parser_class = parser if parser is not None else PydanticConfigParser - self.parser = parser_class(ServiceConfig) - - if service_config is None: - input_args = [] - if config_path: - input_args.append(f"config={config_path}") - if args: - input_args.extend(args) - if kwargs: - input_args.extend([f"{k}={v}" for k, v in kwargs.items()]) - service_config = self.parser.parse_args(*input_args) - - C.service_config = service_config - - if C.service_config.init_logger: - init_logger() - - if llm: - C.update_section_config("llm", **llm) - if embedding_model: - C.update_section_config("embedding_model", **embedding_model) - if vector_store: - C.update_section_config("vector_store", **vector_store) - if token_counter: - C.update_section_config("token_counter", **token_counter) - C.service_config.enable_logo = enable_logo - C.print_logo() - - @staticmethod - def _update_env(key: str, value: str | None): - """Update environment variable if value is provided.""" - if value: - os.environ[key] = value - - @staticmethod - async def start(): - """Initialize the service context and prepare external MCP servers.""" - C.initialize_service_context() - await C.prepare_mcp_servers() - - @staticmethod - def start_sync(): - """Synchronous version of start().""" - C.initialize_service_context() - run_coro_safely(C.prepare_mcp_servers()) - - @staticmethod - async def stop(wait_thread_pool: bool = True, wait_ray: bool = True): - """ - Stop the application and cleanup resources. - - Args: - wait_thread_pool: Whether to wait for thread pool shutdown - wait_ray: Whether to wait for Ray shutdown - """ - await C.close() - C.shutdown_thread_pool(wait=wait_thread_pool) - C.shutdown_ray(wait=wait_ray) - - @staticmethod - def stop_sync(wait_thread_pool: bool = True, wait_ray: bool = True): - """Synchronous version of stop().""" - C.close_sync() - C.shutdown_thread_pool(wait=wait_thread_pool) - C.shutdown_ray(wait=wait_ray) - - async def __aenter__(self): - """Async context manager entry.""" - await self.start() - return self - - def __enter__(self): - """Context manager entry.""" - self.start_sync() - return self - - async def __aexit__(self, exc_type=None, exc_val=None, exc_tb=None): - """Async context manager exit.""" - await self.stop() - return False - - def __exit__(self, exc_type=None, exc_val=None, exc_tb=None): - """Context manager exit.""" - self.stop_sync() - return False - - @staticmethod - async def execute_flow(name: str, **kwargs) -> Response: - """ - Execute a flow asynchronously. - - Args: - name: Name of the flow to execute - **kwargs: Arguments to pass to the flow - - Returns: - Response object from the flow execution - """ - flow: BaseFlow = C.get_flow(name) - return await flow.call(**kwargs) - - @staticmethod - def execute_flow_sync(name: str, **kwargs) -> Response: - """ - Execute a flow synchronously. - - Args: - name: Name of the flow to execute - **kwargs: Arguments to pass to the flow - - Returns: - Response object from the flow execution - """ - flow: BaseFlow = C.get_flow(name) - return flow.call_sync(**kwargs) - - @staticmethod - async def execute_stream_flow(name: str, **kwargs): - """ - Execute a streaming flow asynchronously. - - Args: - name: Name of the streaming flow to execute - **kwargs: Arguments to pass to the flow - - Yields: - Stream chunks from the flow execution - - Raises: - AssertionError: If the flow is not configured for streaming - """ - flow: BaseFlow = C.get_flow(name) - assert flow.stream is True, "non-stream flow is not supported in execute_stream_flow!" - stream_queue = asyncio.Queue() - task = asyncio.create_task(flow.call(stream_queue=stream_queue, **kwargs)) - - async for chunk in execute_stream_task( - stream_queue=stream_queue, - task=task, - task_name=name, - as_bytes=False, - ): - yield chunk - - @staticmethod - def run_service(): - """Run the configured service (HTTP, MCP, or CMD).""" - C.get_service().run() diff --git a/reme_ai/core/context/__init__.py b/reme_ai/core/context/__init__.py deleted file mode 100644 index 7f26d600..00000000 --- a/reme_ai/core/context/__init__.py +++ /dev/null @@ -1,16 +0,0 @@ -"""context""" - -from .base_context import BaseContext -from .prompt_handler import PromptHandler -from .registry import Registry -from .runtime_context import RuntimeContext -from .service_context import ServiceContext, C - -__all__ = [ - "BaseContext", - "PromptHandler", - "Registry", - "RuntimeContext", - "ServiceContext", - "C", -] diff --git a/reme_ai/core/context/base_context.py b/reme_ai/core/context/base_context.py deleted file mode 100644 index dabd8cdb..00000000 --- a/reme_ai/core/context/base_context.py +++ /dev/null @@ -1,41 +0,0 @@ -"""Module providing a dictionary subclass with attribute-style access and pickling support.""" - -from typing import Generic, TypeVar - -_KT = TypeVar("_KT") -_VT = TypeVar("_VT") - - -class BaseContext(dict, Generic[_KT, _VT]): - """A dictionary subclass that enables accessing and modifying keys as attributes.""" - - def __getattr__(self, name: str) -> _VT: - """Retrieve a dictionary item as an attribute.""" - try: - return self[name] - except KeyError as e: - raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'") from e - - def __setattr__(self, name: str, value: _VT) -> None: - """Assign a value to a dictionary item using attribute syntax.""" - self[name] = value - - def __delattr__(self, name: str) -> None: - """Remove a dictionary item using attribute syntax.""" - try: - # Delete item from dict via key - del self[name] - except KeyError as e: - raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'") from e - - def __getstate__(self) -> dict: - """Return the dictionary representation for pickling.""" - return dict(self) - - def __setstate__(self, state: dict) -> None: - """Restore the dictionary state from a pickled object.""" - self.update(state) - - def __reduce__(self): - """Define the reconstruction logic for pickling processes.""" - return self.__class__, (), self.__getstate__() diff --git a/reme_ai/core/context/prompt_handler.py b/reme_ai/core/context/prompt_handler.py deleted file mode 100644 index b428b163..00000000 --- a/reme_ai/core/context/prompt_handler.py +++ /dev/null @@ -1,95 +0,0 @@ -"""Module for managing and formatting prompt templates from files or dictionaries.""" - -from pathlib import Path - -import yaml -from loguru import logger - -from .base_context import BaseContext -from .service_context import C - - -class PromptHandler(BaseContext): - """A context-aware handler for loading, retrieving, and formatting prompt templates.""" - - def __init__(self, language: str = "", **kwargs): - """Initialize the handler with a specific language and optional context data.""" - super().__init__(**kwargs) - self.language: str = language or C.language - - def load_prompt_by_file(self, prompt_file_path: Path | str = None): - """Load prompt configurations from a YAML file into the context.""" - if prompt_file_path is None: - return self - - if isinstance(prompt_file_path, str): - prompt_file_path = Path(prompt_file_path) - - if not prompt_file_path.exists(): - return self - - with prompt_file_path.open(encoding="utf-8") as f: - # Load YAML content using the full loader - prompt_dict = yaml.load(f, yaml.FullLoader) - self.load_prompt_dict(prompt_dict) - return self - - def load_prompt_dict(self, prompt_dict: dict = None): - """Merge a dictionary of prompt strings into the current context.""" - if not prompt_dict: - return self - - for key, value in prompt_dict.items(): - if isinstance(value, str): - if key in self: - logger.warning(f"Overwriting prompt key={key}, old_value={self[key]}, new_value={value}") - else: - logger.debug(f"Adding new prompt key={key}, value={value}") - self[key] = value - return self - - def get_prompt(self, prompt_name: str): - """Retrieve a prompt by name, automatically appending the language suffix if needed.""" - key: str = prompt_name - if self.language and not key.endswith(self.language.strip()): - key += "_" + self.language.strip() - - assert key in self, f"prompt_name={key} not found." - return self[key].strip() - - def prompt_format(self, prompt_name: str, **kwargs) -> str: - """Format a prompt by filtering flagged lines and filling template variables.""" - prompt = self.get_prompt(prompt_name) - - # Separate boolean flags from string formatting arguments - flag_kwargs = {k: v for k, v in kwargs.items() if isinstance(v, bool)} - other_kwargs = {k: v for k, v in kwargs.items() if not isinstance(v, bool)} - - if flag_kwargs: - split_prompt = [] - for line in prompt.strip().split("\n"): - hit = False - hit_flag = True - for key, flag in flag_kwargs.items(): - if not line.startswith(f"[{key}]"): - continue - - hit = True - hit_flag = flag - # Remove the flag prefix from the line - line = line.strip(f"[{key}]") - break - - # Include line if no flag is present or if the flag evaluates to True - if not hit: - split_prompt.append(line) - elif hit_flag: - split_prompt.append(line) - - prompt = "\n".join(split_prompt) - - if other_kwargs: - # Apply standard Python string formatting - prompt = prompt.format(**other_kwargs) - - return prompt diff --git a/reme_ai/core/context/registry.py b/reme_ai/core/context/registry.py deleted file mode 100644 index 08ea2271..00000000 --- a/reme_ai/core/context/registry.py +++ /dev/null @@ -1,46 +0,0 @@ -"""Module providing a registry class for managing class-to-name mappings via decorators.""" - -import inspect -from typing import Callable, TypeVar - -from .base_context import BaseContext - -T = TypeVar('T') - - -class Registry(BaseContext): - """A registry container that uses decorators to map and store class references.""" - - def register(self, name: str | type = "", add_cls: bool = True) -> Callable[[type[T]], type[T]] | type[T]: - """Return a decorator that registers a class under a specific name in the registry. - - Can be used in three ways: - - @C.register_op() # with empty parentheses, uses class name - - @C.register_op # without parentheses, uses class name - - @C.register_op("custom_name") # with custom name - - Args: - name: Either a string name for the class, or the class itself when used without parentheses - add_cls: Whether to actually add the class to the registry - - Returns: - Either a decorator function or the registered class itself - """ - - def decorator(cls): - if add_cls: - # Use provided name or default to the class name as the key - key = name if isinstance(name, str) and name else cls.__name__ - self[key] = cls - return cls - - # If used without parentheses: @C.register_op - if inspect.isclass(name): - cls = name - # Register with class name as key - if add_cls: - self[cls.__name__] = cls - return cls - - # If used with parentheses: @C.register_op() or @C.register_op("name") - return decorator diff --git a/reme_ai/core/context/service_context.py b/reme_ai/core/context/service_context.py deleted file mode 100644 index e22147e9..00000000 --- a/reme_ai/core/context/service_context.py +++ /dev/null @@ -1,537 +0,0 @@ -"""Module for managing global service configurations and component registries via a singleton context.""" - -from concurrent.futures import ThreadPoolExecutor -from typing import TYPE_CHECKING - -from loguru import logger - -from .base_context import BaseContext -from .registry import Registry -from ..enumeration import RegistryEnum -from ..schema import ServiceConfig -from ..utils import singleton, print_logo - -if TYPE_CHECKING: - from ..llm import BaseLLM - from ..embedding import BaseEmbeddingModel - from ..vector_store import BaseVectorStore - from ..token_counter import BaseTokenCounter - from ..flow import BaseFlow - from ..service import BaseService - - -@singleton -class ServiceContext(BaseContext): - """A singleton container for global application state, thread pools, and component registries. - - This class serves as the central management hub for the entire ReMe application, providing: - - Service configuration management - - Component registration and instantiation (LLMs, embeddings, vector stores, etc.) - - Thread pool and Ray distributed computing management - - MCP (Model Context Protocol) server integration - - The singleton pattern ensures only one instance exists throughout the application lifecycle, - accessible via the global `C` variable exported at the bottom of this module. - """ - - def __init__(self, **kwargs): - """Initialize the global context with configuration objects and specialized registries. - - Sets up: - - Empty service configuration placeholder - - Thread pool for concurrent operations - - Registry dictionaries for class registration (templates) - - Instance dictionaries for instantiated objects (actual instances) - - MCP server mapping for external tool integration - """ - super().__init__(**kwargs) - - # Service configuration and runtime settings - self.service_config: ServiceConfig | None = None - self.language: str = "" - self.thread_pool: ThreadPoolExecutor | None = None - - # Registry system: stores class definitions for different component types - self.registry_dict: dict[RegistryEnum, Registry] = {v: Registry() for v in RegistryEnum.__members__.values()} - - # Instance system: stores instantiated objects created from registered classes - self.instance_dict: dict[RegistryEnum, dict] = {v: {} for v in RegistryEnum.__members__.values()} - - # MCP server mapping: maps server_name -> {tool_name: ToolCall} - self.mcp_server_mapping: dict[str, dict] = {} - - # Initialization flag: ensures initialize_service_context is called only once - self._initialized: bool = False - - def register(self, name: str, register_type: RegistryEnum): - """Return a decorator to register a component within a specific registry category. - - Args: - name: The registration name for the component (used for lookup) - register_type: The type of registry (LLM, EMBEDDING_MODEL, VECTOR_STORE, etc.) - - Returns: - A decorator function that registers the decorated class - - Example: - @C.register("my_llm", RegistryEnum.LLM) - class MyLLM(BaseLLM): - pass - """ - return self.registry_dict[register_type].register(name=name) - - def register_llm(self, name: str = ""): - """Register a Large Language Model class.""" - return self.register(name=name, register_type=RegistryEnum.LLM) - - def register_embedding_model(self, name: str = ""): - """Register an embedding model class.""" - return self.register(name=name, register_type=RegistryEnum.EMBEDDING_MODEL) - - def register_vector_store(self, name: str = ""): - """Register a vector store implementation class.""" - return self.register(name=name, register_type=RegistryEnum.VECTOR_STORE) - - def register_op(self, name: str = ""): - """Register an operation (Op) class.""" - return self.register(name=name, register_type=RegistryEnum.OP) - - def register_flow(self, name: str = ""): - """Register a workflow or logic flow class.""" - return self.register(name=name, register_type=RegistryEnum.FLOW) - - def register_service(self, name: str = ""): - """Register a backend service class.""" - return self.register(name=name, register_type=RegistryEnum.SERVICE) - - def register_token_counter(self, name: str = ""): - """Register a token counting utility class.""" - return self.register(name=name, register_type=RegistryEnum.TOKEN_COUNTER) - - def get_model_class(self, name: str, register_type: RegistryEnum): - """Retrieve a registered class by name from a specific registry category. - - Args: - name: The registration name of the class - register_type: The type of registry to search in - - Returns: - The registered class (not an instance, but the class itself) - - Raises: - AssertionError: If the class is not found in the registry - """ - assert name in self.registry_dict[register_type], f"{name} not in registry_dict[{register_type}]" - return self.registry_dict[register_type][name] - - def get_llm_class(self, name: str): - """Get the LLM class registered under the given name.""" - return self.get_model_class(name, RegistryEnum.LLM) - - def get_embedding_model_class(self, name: str): - """Get the embedding model class registered under the given name.""" - return self.get_model_class(name, RegistryEnum.EMBEDDING_MODEL) - - def get_vector_store_class(self, name: str): - """Get the vector store class registered under the given name.""" - return self.get_model_class(name, RegistryEnum.VECTOR_STORE) - - def get_op_class(self, name: str): - """Get the operation class registered under the given name.""" - return self.get_model_class(name, RegistryEnum.OP) - - def get_flow_class(self, name: str): - """Get the flow class registered under the given name.""" - return self.get_model_class(name, RegistryEnum.FLOW) - - def get_service_class(self, name: str): - """Get the service class registered under the given name.""" - return self.get_model_class(name, RegistryEnum.SERVICE) - - def get_token_counter_class(self, name: str): - """Get the token counter class registered under the given name.""" - return self.get_model_class(name, RegistryEnum.TOKEN_COUNTER) - - def get_llm(self, name: str) -> "BaseLLM": - """Retrieve a specific LLM instance by name. - - Args: - name: The name of the LLM instance (typically 'default' or custom name) - - Returns: - The instantiated LLM object - - Raises: - KeyError: If no LLM with the given name exists - """ - return self.instance_dict[RegistryEnum.LLM][name] - - def get_embedding_model(self, name: str) -> "BaseEmbeddingModel": - """Retrieve a specific embedding model instance by name. - - Args: - name: The name of the embedding model instance (typically 'default') - - Returns: - The instantiated embedding model object - - Raises: - KeyError: If no embedding model with the given name exists - """ - return self.instance_dict[RegistryEnum.EMBEDDING_MODEL][name] - - def get_vector_store(self, name: str) -> "BaseVectorStore": - """Retrieve a specific vector store instance by name. - - Args: - name: The name of the vector store instance (typically 'default') - - Returns: - The instantiated vector store object - - Raises: - KeyError: If no vector store with the given name exists - """ - return self.instance_dict[RegistryEnum.VECTOR_STORE][name] - - def get_token_counter(self, name: str) -> "BaseTokenCounter": - """Retrieve a specific token counter instance by name. - - Args: - name: The name of the token counter instance (typically 'default') - - Returns: - The instantiated token counter object - - Raises: - KeyError: If no token counter with the given name exists - """ - return self.instance_dict[RegistryEnum.TOKEN_COUNTER][name] - - def get_flow(self, name: str) -> "BaseFlow": - """Retrieve a specific flow instance by name. - - Args: - name: The name of the flow instance - - Returns: - The instantiated flow object - - Raises: - KeyError: If no flow with the given name exists - """ - return self.instance_dict[RegistryEnum.FLOW][name] - - def get_service(self) -> "BaseService": - """Retrieve the default service instance. - - Returns: - The instantiated service backend (HTTP, MCP, or CMD service) - - Raises: - KeyError: If the default service was not initialized - """ - return self.instance_dict[RegistryEnum.SERVICE]["default"] - - def update_section_config(self, section_name: str, **kwargs): - """Update a specific section of the service config with new values. - - Args: - section_name: Name of the config section (e.g., 'llm', 'embedding_model') - **kwargs: Key-value pairs to update in the default configuration - - Raises: - KeyError: If the default config for the section doesn't exist - - Example: - update_section_config('llm', temperature=0.8, max_tokens=1000) - """ - if not hasattr(self.service_config, section_name) or not kwargs: - return - - section_dict: dict = getattr(self.service_config, section_name) - if "default" not in section_dict: - raise KeyError(f"Default `{section_name}` config not found") - - current_config = section_dict["default"] - section_dict["default"] = current_config.model_copy(update=kwargs, deep=True) - - def initialize_service_context(self): - """Initialize the service context with the configuration. - - This is the main initialization method that sets up all system components in order: - 1. Language settings - 2. Thread pool for concurrent operations - 3. Ray cluster (if configured for distributed computing) - 4. LLM instances - 5. Embedding model instances - 6. Token counter instances - 7. Vector store instances (with their embedding models) - 8. Flow instances (both registered and configured) - 9. Service backend instance - - Note: This method should be called after service_config is set. - This method can only be called once. Subsequent calls will be ignored. - """ - if self._initialized: - logger.warning("initialize_service_context has already been called. Skipping re-initialization.") - return - - self.language = self.service_config.language - self.thread_pool = ThreadPoolExecutor(max_workers=self.service_config.thread_pool_max_workers) - - # Initialize Ray for distributed computing if configured - if self.service_config.ray_max_workers > 1: - import ray - - ray.init(num_cpus=self.service_config.ray_max_workers) - - # Initialize components in dependency order - self._initialize_llm() - self._initialize_embedding_model() - self._initialize_token_counter() - self._initialize_vector_store() # Depends on embedding models - self._initialize_flow() - self._initialize_service() - - # Mark as initialized - self._initialized = True - - def _initialize_llm(self): - """Initialize all configured LLM instances. - - For each LLM configuration: - - Retrieves the corresponding registered LLM class by backend name - - Instantiates it with model_name and additional configuration - - Stores the instance in instance_dict for later retrieval - """ - for name, config in self.service_config.llm.items(): - llm_cls = self.get_llm_class(config.backend) - self.instance_dict[RegistryEnum.LLM][name] = llm_cls(model_name=config.model_name, **config.model_extra) - - def _initialize_embedding_model(self): - """Initialize all configured embedding model instances. - - For each embedding model configuration: - - Retrieves the corresponding registered embedding model class by backend name - - Instantiates it with model_name and additional configuration - - Stores the instance in instance_dict for later retrieval - """ - for name, config in self.service_config.embedding_model.items(): - embedding_model_cls = self.get_embedding_model_class(config.backend) - self.instance_dict[RegistryEnum.EMBEDDING_MODEL][name] = embedding_model_cls( - model_name=config.model_name, - **config.model_extra, - ) - - def _initialize_token_counter(self): - """Initialize all configured token counter instances. - - For each token counter configuration: - - Retrieves the corresponding registered token counter class by backend name - - Instantiates it with model_name and additional configuration - - Stores the instance in instance_dict for later retrieval - """ - for name, config in self.service_config.token_counter.items(): - token_counter_cls = self.get_token_counter_class(config.backend) - self.instance_dict[RegistryEnum.TOKEN_COUNTER][name] = token_counter_cls( - model_name=config.model_name, - **config.model_extra, - ) - - def _initialize_vector_store(self): - """Initialize all configured vector stores with their embedding models. - - For each vector store configuration: - - Retrieves the corresponding registered vector store class by backend name - - Retrieves the associated embedding model instance by name - - Instantiates the vector store with collection name, embedding model, and extra config - - Stores the instance in instance_dict for later retrieval - - Note: This must be called after _initialize_embedding_model() since vector stores - depend on embedding model instances. - """ - for name, config in self.service_config.vector_store.items(): - vector_store_cls = self.get_vector_store_class(config.backend) - self.instance_dict[RegistryEnum.VECTOR_STORE][name] = vector_store_cls( - collection_name=config.collection_name, - embedding_model=self.instance_dict[RegistryEnum.EMBEDDING_MODEL][config.embedding_model], - **config.model_extra, - ) - - def _filter_flows(self, name: str) -> bool: - """Filter flows based on enabled_flows and disabled_flows configuration. - - The filtering logic follows this priority: - 1. If enabled_flows is set: only flows in the list are loaded - 2. Else if disabled_flows is set: all flows except those in the list are loaded - 3. Otherwise: all flows are loaded - - Args: - name: The flow name to check - - Returns: - True if the flow should be loaded, False otherwise - """ - if self.service_config.enabled_flows: - return name in self.service_config.enabled_flows - elif self.service_config.disabled_flows: - return name not in self.service_config.disabled_flows - else: - return True - - def _initialize_flow(self): - """Initialize all flows from both registry and configuration. - - Flows can be defined in two ways: - 1. Registered flows: Python classes decorated with @register_flow - 2. Configuration flows: Defined in config as ExpressionFlow instances - - Process: - 1. First, instantiate all registered flow classes (from decorators) - - Filter based on enabled_flows/disabled_flows - - Create instance with the flow name - - 2. Then, instantiate all configured flows (from config file) - - Filter based on enabled_flows/disabled_flows - - Create ExpressionFlow instances with flow configuration - - Note: Configuration flows can override registered flows with the same name. - """ - - # Initialize flows from registry (decorator-based registration) - for name, flow_cls in self.registry_dict[RegistryEnum.FLOW].items(): - if not self._filter_flows(name): - continue - flow: "BaseFlow" = flow_cls(name=name) - self.instance_dict[RegistryEnum.FLOW][flow.name] = flow - - # Initialize flows from configuration (config-based definition) - from ..flow import ExpressionFlow - - for name, flow_config in self.service_config.flow.items(): - if not self._filter_flows(name): - continue - flow_config.name = name - flow: BaseFlow = ExpressionFlow(flow_config=flow_config) - self.instance_dict[RegistryEnum.FLOW][name] = flow - - def _initialize_service(self): - """Initialize the service backend instance. - - Creates an instance of the configured service backend (e.g., HTTP, MCP, or CMD service) - and stores it in the instance dictionary under the 'default' key. - """ - service_cls = self.get_service_class(self.service_config.backend) - self.instance_dict[RegistryEnum.SERVICE]["default"] = service_cls() - - async def prepare_mcp_servers(self): - """Prepare and initialize MCP (Model Context Protocol) server connections. - - This method: - 1. Checks if MCP servers are configured - 2. Creates an MCP client instance - 3. For each configured server: - - Lists available tool calls from the server - - Builds a mapping of tool_name -> ToolCall object - - Logs available tools for debugging - - The mcp_server_mapping is structured as: - { - "server_name": { - "tool_name": ToolCall(...), - ... - }, - ... - } - - This allows the application to discover and use external tools provided by MCP servers. - """ - if not self.service_config.mcp_servers: - return - - from ..utils import MCPClient - - mcp_client = MCPClient(config={"mcpServers": self.service_config.mcp_servers}) - for server_name in self.service_config.mcp_servers.keys(): - try: - # Retrieve all available tool calls from this MCP server - tool_calls = await mcp_client.list_tool_calls(server_name=server_name, return_dict=False) - - # Build mapping: tool_name -> ToolCall for quick lookup - self.mcp_server_mapping[server_name] = {tool_call.name: tool_call for tool_call in tool_calls} - - # Log discovered tools for debugging - for tool_call in tool_calls: - logger.info(f"list_tool_calls: {server_name}@{tool_call.name} {tool_call.simple_input_dump()}") - - except Exception as e: - logger.exception(f"list_tool_calls: {server_name} error: {e}") - - def print_logo(self): - """Print the ReMe logo if enabled in configuration.""" - if self.service_config.enable_logo: - print_logo(service_config=self.service_config) - - async def close(self): - """Close all service components asynchronously. - - Gracefully closes all instantiated components in order: - 1. Vector stores (closes database connections) - 2. LLMs (closes API clients and connections) - 3. Embedding models (closes API clients and connections) - - This method should be called when shutting down the application - to ensure all resources are properly released. - """ - for _, vector_store in self.instance_dict[RegistryEnum.VECTOR_STORE].items(): - await vector_store.close() - - for _, llm in self.instance_dict[RegistryEnum.LLM].items(): - await llm.close() - - for _, embedding_model in self.instance_dict[RegistryEnum.EMBEDDING_MODEL].items(): - await embedding_model.close() - - def close_sync(self): - """Close all service components synchronously. - - Synchronous version of close() for non-async contexts. - Closes LLMs and embedding models without using async/await. - - Note: Vector stores are not closed here as they typically require async operations. - """ - for _, llm in self.instance_dict[RegistryEnum.LLM].items(): - llm.close_sync() - - for _, embedding_model in self.instance_dict[RegistryEnum.EMBEDDING_MODEL].items(): - embedding_model.close_sync() - - def shutdown_thread_pool(self, wait: bool = True): - """Shutdown the thread pool executor. - - Args: - wait: If True, blocks until all pending futures are executed. - If False, returns immediately and pending futures may be cancelled. - """ - if self.thread_pool: - self.thread_pool.shutdown(wait=wait) - - def shutdown_ray(self, wait: bool = True): - """Shutdown Ray cluster if it was initialized. - - Args: - wait: If True, waits for Ray to fully shutdown. - If False, returns immediately without waiting. - - Note: Only shuts down Ray if it was configured with ray_max_workers > 1. - """ - if self.service_config and self.service_config.ray_max_workers > 1: - import ray - - ray.shutdown(_exiting_interpreter=not wait) - - -# Export a global singleton instance for easy access across the application -# This is the primary way to access the service context throughout the codebase -C = ServiceContext() diff --git a/reme_ai/core/enumeration/__init__.py b/reme_ai/core/enumeration/__init__.py deleted file mode 100644 index 7202a949..00000000 --- a/reme_ai/core/enumeration/__init__.py +++ /dev/null @@ -1,17 +0,0 @@ -"""enumeration""" - -from .chunk_enum import ChunkEnum -from .http_enum import HttpEnum -from .json_schema_enum import JsonSchemaEnum -from .memory_type import MemoryType -from .registry_enum import RegistryEnum -from .role import Role - -__all__ = [ - "ChunkEnum", - "HttpEnum", - "JsonSchemaEnum", - "MemoryType", - "RegistryEnum", - "Role", -] diff --git a/reme_ai/core/enumeration/chunk_enum.py b/reme_ai/core/enumeration/chunk_enum.py deleted file mode 100644 index dbe37106..00000000 --- a/reme_ai/core/enumeration/chunk_enum.py +++ /dev/null @@ -1,25 +0,0 @@ -"""Defines the types of data chunks used in streaming responses.""" - -from enum import Enum - - -class ChunkEnum(str, Enum): - """Enumeration of possible chunk categories for stream processing.""" - - # Internal reasoning or chain-of-thought process - THINK = "think" - - # The final generated response content - ANSWER = "answer" - - # Metadata or calls related to external tools - TOOL = "tool" - - # Resource consumption and token usage statistics - USAGE = "usage" - - # Error messages or exception details - ERROR = "error" - - # Final signal indicating the completion of the stream - DONE = "done" diff --git a/reme_ai/core/enumeration/http_enum.py b/reme_ai/core/enumeration/http_enum.py deleted file mode 100644 index 19622242..00000000 --- a/reme_ai/core/enumeration/http_enum.py +++ /dev/null @@ -1,22 +0,0 @@ -"""Provides a collection of standard HTTP request methods.""" - -from enum import Enum - - -class HttpEnum(str, Enum): - """Enumeration of supported HTTP methods for network requests.""" - - # Retrieves data from a specified resource - GET = "get" - - # Submits data to be processed to a specified resource - POST = "post" - - # Identical to GET but only retrieves the response headers - HEAD = "head" - - # Uploads or replaces the representation of a target resource - PUT = "put" - - # Deletes the specified resource from the server - DELETE = "delete" diff --git a/reme_ai/core/enumeration/json_schema_enum.py b/reme_ai/core/enumeration/json_schema_enum.py deleted file mode 100644 index 507645f4..00000000 --- a/reme_ai/core/enumeration/json_schema_enum.py +++ /dev/null @@ -1,18 +0,0 @@ -"""Defines the standard data types supported by JSON Schema.""" - -from enum import Enum - - -class JsonSchemaEnum(Enum): - """Enumeration of valid JSON Schema data types.""" - - STRING = str - NUMBER = float - INTEGER = int - OBJECT = dict - ARRAY = list - BOOLEAN = bool - - def __str__(self) -> str: - """Returns the string representation of the enum value.""" - return self.name.lower() diff --git a/reme_ai/core/enumeration/memory_type.py b/reme_ai/core/enumeration/memory_type.py deleted file mode 100644 index 22d35481..00000000 --- a/reme_ai/core/enumeration/memory_type.py +++ /dev/null @@ -1,25 +0,0 @@ -"""Memory type enumeration for the three-layer memory architecture.""" - -from enum import Enum - - -class MemoryType(str, Enum): - """ - Three-layer memory architecture for agent memory management. - - Layer 1 - High-level Abstraction Memory: - - IDENTITY: Self-cognition (identity, personality, current state) - - PERSONAL: Person-specific memory (preferences and context about specific individuals) - - PROCEDURAL: Procedural memory (how-to knowledge, e.g., 4 steps to write financial reports) - - TOOL: Tool memory (tool usage patterns, success rates, token consumption, latency) - - Layer 2 - Summary Memory (Compressed): Summarized digest of raw message history - Layer 3 - History Memory (Raw): Raw message history - """ - - IDENTITY = "identity" - PERSONAL = "personal" - PROCEDURAL = "procedural" - TOOL = "tool" - SUMMARY = "summary" - HISTORY = "history" diff --git a/reme_ai/core/enumeration/registry_enum.py b/reme_ai/core/enumeration/registry_enum.py deleted file mode 100644 index 876c06b8..00000000 --- a/reme_ai/core/enumeration/registry_enum.py +++ /dev/null @@ -1,28 +0,0 @@ -"""Defines the registry categories for core components of the system.""" - -from enum import Enum - - -class RegistryEnum(str, Enum): - """Enumeration of component types registered within the application lifecycle.""" - - # Large Language Model interfaces - LLM = "llm" - - # Models used for generating vector embeddings - EMBEDDING_MODEL = "embedding_model" - - # Databases or storage systems for vector search - VECTOR_STORE = "vector_store" - - # Atomic operations or functional units - OP = "op" - - # Orchestrated sequences of operations or workflows - FLOW = "flow" - - # External APIs or shared internal services - SERVICE = "service" - - # Utilities for tracking and limiting token consumption - TOKEN_COUNTER = "token_counter" diff --git a/reme_ai/core/enumeration/role.py b/reme_ai/core/enumeration/role.py deleted file mode 100644 index 4acad7e5..00000000 --- a/reme_ai/core/enumeration/role.py +++ /dev/null @@ -1,19 +0,0 @@ -"""Defines the participant roles in a chat completion sequence.""" - -from enum import Enum - - -class Role(str, Enum): - """Enumeration of standard personas involved in a conversation flow.""" - - # High-level instructions to guide the model's behavior - SYSTEM = "system" - - # Input or queries provided by the human user - USER = "user" - - # Responses or messages generated by the AI model - ASSISTANT = "assistant" - - # Output or results returned from external tool executions - TOOL = "tool" diff --git a/reme_ai/core/flow/simple_flow.py b/reme_ai/core/flow/simple_flow.py deleted file mode 100644 index 51a1381d..00000000 --- a/reme_ai/core/flow/simple_flow.py +++ /dev/null @@ -1,17 +0,0 @@ -"""Simple flow implementation that directly uses a predefined flow operation.""" - -from .base_flow import BaseFlow -from ..op import BaseOp -from ..schema import ToolCall - - -class SimpleFlow(BaseFlow): - """Simple flow that directly uses a predefined flow operation.""" - - def _build_flow(self) -> BaseOp: - assert self._flow_op is not None - return self._flow_op.copy() - - def _build_tool_call(self) -> ToolCall: - assert self._flow_op is not None - return self._flow_op.tool_call diff --git a/reme_ai/core/main.py b/reme_ai/core/main.py deleted file mode 100644 index 81811c1f..00000000 --- a/reme_ai/core/main.py +++ /dev/null @@ -1,48 +0,0 @@ -"""ReMe application classes for simplified configuration and execution.""" - -import sys - -from .application import Application -from .config import ReMeConfigParser - - -class ReMeApp(Application): - """ReMe application with config file support and flow execution methods.""" - - def __init__( - self, - *args, - llm_api_key: str | None = None, - llm_api_base: str | None = None, - embedding_api_key: str | None = None, - embedding_api_base: str | None = None, - config_path: str | None = None, - enable_logo: bool = True, - **kwargs, - ): - super().__init__( - *args, - llm_api_key=llm_api_key, - llm_api_base=llm_api_base, - embedding_api_key=embedding_api_key, - embedding_api_base=embedding_api_base, - service_config=None, - parser=ReMeConfigParser, - config_path=config_path, - enable_logo=enable_logo, - **kwargs, - ) - - async def async_execute(self, name: str, **kwargs) -> dict: - """Execute a flow asynchronously and return the result as a dictionary.""" - return (await self.execute_flow(name=name, **kwargs)).model_dump() - - -def main(): - """Main entry point for running ReMe application from command line.""" - with ReMeApp(*sys.argv[1:]) as app: - app.run_service() - - -if __name__ == "__main__": - main() diff --git a/reme_ai/core/schema/__init__.py b/reme_ai/core/schema/__init__.py deleted file mode 100644 index b7b73719..00000000 --- a/reme_ai/core/schema/__init__.py +++ /dev/null @@ -1,42 +0,0 @@ -"""schema""" - -from .memory_node import MemoryNode -from .message import ContentBlock, Message, Trajectory -from .request import Request -from .response import Response -from .service_config import ( - CmdConfig, - EmbeddingModelConfig, - FlowConfig, - HttpConfig, - LLMConfig, - MCPConfig, - ServiceConfig, - TokenCounterConfig, - VectorStoreConfig, -) -from .stream_chunk import StreamChunk -from .tool_call import ToolAttr, ToolCall -from .vector_node import VectorNode - -__all__ = [ - "MemoryNode", - "ContentBlock", - "EmbeddingModelConfig", - "FlowConfig", - "HttpConfig", - "LLMConfig", - "MCPConfig", - "Message", - "Request", - "Response", - "ServiceConfig", - "StreamChunk", - "TokenCounterConfig", - "Trajectory", - "ToolAttr", - "ToolCall", - "VectorNode", - "VectorStoreConfig", - "CmdConfig", -] diff --git a/reme_ai/core/schema/memory_node.py b/reme_ai/core/schema/memory_node.py deleted file mode 100644 index 304dfe5c..00000000 --- a/reme_ai/core/schema/memory_node.py +++ /dev/null @@ -1,223 +0,0 @@ -"""Memory schema module for the ReMe AI system. - -This module defines the MemoryNode class for storing and retrieving -memories in the ReMe system. -""" - -import datetime -import hashlib -import json -from typing import Any - -from pydantic import BaseModel, Field, model_validator - -from .vector_node import VectorNode -from ..enumeration import MemoryType - - -def get_now_time() -> str: - """Get current timestamp in YYYY-MM-DD HH:MM:SS format. - - Returns: - str: Current timestamp string in format 'YYYY-MM-DD HH:MM:SS'. - """ - return datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S") - - -# Length of the memory ID (first N characters of SHA-256 hash) -MEMORY_ID_LENGTH: int = 16 - - -class MemoryNode(BaseModel): - """Memory node for storing memories in the ReMe system. - - Attributes: - memory_id: Unique identifier, auto-generated from content hash. - memory_type: Type of memory (e.g., SUMMARY, PERSONAL). - memory_target: Target or topic this memory relates to. - when_to_use: Condition description for vector retrieval. - content: Actual memory content. - ref_memory_id: Reference to related raw history memory. - time_created: Creation timestamp. - time_modified: Last modification timestamp. - author: Author or source of this memory. - score: Relevance or importance score. - metadata: Additional metadata for extensibility. - """ - - memory_id: str = Field(default="", description="Unique memory identifier") - memory_type: MemoryType = Field(default=..., description="Type of memory") - memory_target: str = Field(default="", description="Target or topic of the memory") - when_to_use: str = Field(default="", description="Condition description for vector retrieval") - content: str = Field(default="", description="Actual memory content") - ref_memory_id: str = Field(default="", description="Reference to related raw history memory ID") - - time_created: str = Field(default_factory=get_now_time, description="Creation timestamp") - time_modified: str = Field(default_factory=get_now_time, description="Last modification timestamp") - author: str = Field(default="", description="Author or source of the memory") - score: float = Field(default=0, description="Relevance or importance score") - - metadata: dict[str, Any] = Field(default_factory=dict, description="Additional metadata") - - def _update_modified_time(self) -> "MemoryNode": - """Update time_modified to current timestamp. - - Returns: - Self: Returns self for method chaining. - """ - self.time_modified = get_now_time() - return self - - def _update_memory_id(self) -> "MemoryNode": - """Generate memory_id from SHA-256 hash of content. - - Takes the first MEMORY_ID_LENGTH characters of the hash. - - Returns: - Self: Returns self for method chaining. - """ - if not self.content: - return self - - hash_obj = hashlib.sha256(self.content.encode("utf-8")) - hex_dig = hash_obj.hexdigest() - self.memory_id = hex_dig[:MEMORY_ID_LENGTH] - return self - - @model_validator(mode="after") - def _update_after_init(self) -> "MemoryNode": - """Post-initialization validator. - - Auto-generates memory_id from content if not provided. - - Returns: - Self: Returns self for method chaining. - """ - if not self.memory_id: - self._update_memory_id() - return self - - def __setattr__(self, name: str, value): - """Auto-update timestamps and memory_id when content or when_to_use changes. - - Args: - name: Attribute name being set. - value: New value for the attribute. - """ - should_update: bool = name in ("when_to_use", "content") and getattr(self, name, None) != value - super().__setattr__(name, value) - if should_update: - self._update_modified_time() - if name == "content": - self._update_memory_id() - - def to_vector_node(self) -> VectorNode: - """Convert to VectorNode for vector storage. - - When when_to_use is set, use it as vector content and store content in metadata. - When when_to_use is empty, use content as vector content directly. - - Returns: - VectorNode: Vector node representation of this memory. - """ - # Build base metadata (shared fields) - metadata: dict[str, Any] = { - "memory_type": self.memory_type.value, - "memory_target": self.memory_target, - "ref_memory_id": self.ref_memory_id, - "time_created": self.time_created, - "time_modified": self.time_modified, - "author": self.author, - "score": self.score, - **self.metadata, - } - - if self.when_to_use: - # Use when_to_use for vector embedding, store content in metadata - vector_content = self.when_to_use - metadata["content"] = self.content - else: - # Use content directly for vector embedding - vector_content = self.content - - return VectorNode( - vector_id=self.memory_id, - content=vector_content, - metadata=metadata, - ) - - def format_memory(self) -> str: - """Format memory as human-readable string. - - Returns: - str: Formatted string with when_to_use, content, and ref_memory_id. - """ - parts: list[str] = [ - f"memory_id={self.memory_id}", - ] - - if self.when_to_use: - parts.append(self.when_to_use) - - if self.content: - parts.append(self.content) - - if self.metadata: - parts.append(f"metadata={json.dumps(self.metadata, ensure_ascii=False)}") - - if self.ref_memory_id: - parts.append(f"ref_memory_id={self.ref_memory_id}") - - return " ".join(parts) - - @classmethod - def from_vector_node(cls, node: VectorNode) -> "MemoryNode": - """Reconstruct MemoryNode from VectorNode. - - Reverses the to_vector_node conversion: - - If metadata contains 'content': node.content -> when_to_use, metadata['content'] -> content - - Otherwise: node.content -> content, when_to_use remains empty - - Args: - node: VectorNode containing memory data. - - Returns: - Self: Reconstructed MemoryNode instance. - - Raises: - ValueError: If memory_type in metadata is invalid. - """ - metadata = node.metadata.copy() - memory_type_str = metadata.pop("memory_type", None) - - try: - memory_type: MemoryType = MemoryType(memory_type_str) - except ValueError as e: - raise ValueError( - f"Invalid memory_type '{memory_type_str}' in VectorNode metadata. " - f"Valid types are: {[t.value for t in MemoryType]}", - ) from e - - # Restore when_to_use and content based on metadata structure - if "content" in metadata: - # Original had when_to_use set - when_to_use = node.content - content = metadata.pop("content", "") - else: - # Original had empty when_to_use - when_to_use = "" - content = node.content - - return cls( - memory_id=node.vector_id, - memory_type=memory_type, - memory_target=metadata.pop("memory_target", ""), - when_to_use=when_to_use, - content=content, - ref_memory_id=metadata.pop("ref_memory_id", ""), - time_created=metadata.pop("time_created", ""), - time_modified=metadata.pop("time_modified", ""), - author=metadata.pop("author", ""), - score=metadata.pop("score", 0), - metadata=metadata, - ) diff --git a/reme_ai/core/schema/message.py b/reme_ai/core/schema/message.py deleted file mode 100644 index 6c3299e7..00000000 --- a/reme_ai/core/schema/message.py +++ /dev/null @@ -1,164 +0,0 @@ -"""Data models for multi-modal conversation history and LLM interaction trajectories.""" - -import datetime -import json -import re - -from pydantic import BaseModel, ConfigDict, Field, model_validator - -from .tool_call import ToolCall -from ..enumeration import Role - - -class ContentBlock(BaseModel): - """ - Individual unit of multi-modal content like text, images, or video. - examples: - { - "type": "image_url", - "image_url": { - "url": "https://img.alicdn.com/imgextra/i1/O1CN01gDEY8M1W114Hi3XcN_!!6000000002727-0-tps-1024-406.jpg" - }, - } - - { - "type": "video", - "video": [ - "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/xzsgiz/football1.jpg", - "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/tdescd/football2.jpg", - "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/zefdja/football3.jpg", - "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/aedbqh/football4.jpg", - ], - } - - { - "type": "text", - "text": "How do you solve this problem?" - } - """ - - model_config = ConfigDict(extra="allow") - - type: str = Field(default="") - content: str | dict | list = Field(default="") - - @model_validator(mode="before") - @classmethod - def init_block(cls, data: dict) -> dict: - """Dynamically maps the type-specific key to the content field.""" - content_type = data.get("type", "") - if content_type and content_type in data: - data["content"] = data[content_type] - return data - - def simple_dump(self) -> dict: - """Serializes the block into an API-compatible dictionary format.""" - return { - "type": self.type, - self.type: self.content, - **self.model_extra, - } - - -class Message(BaseModel): - """Data model for a single dialogue entry including roles and tool interactions.""" - - name: str | None = Field(default=None) - role: Role = Field(default=Role.USER) - content: str | list[ContentBlock] = Field(default="") - reasoning_content: str = Field(default="") - tool_calls: list[ToolCall] = Field(default_factory=list) - tool_call_id: str = Field(default="") - time_created: str = Field(default_factory=lambda: datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")) - metadata: dict = Field(default_factory=dict) - - def dump_content(self) -> str | list[dict]: - """Returns content as a raw string or a list of serialized blocks.""" - if isinstance(self.content, str): - return self.content - return [block.simple_dump() for block in self.content] - - def simple_dump( - self, - add_name: bool = False, - add_reasoning: bool = True, - add_time_created: bool = False, - add_metadata: bool = False, - enable_json_dump: bool = False, - ) -> dict | str: - """Transforms the message into a simplified dictionary for standard APIs.""" - result = {} - if add_name and self.name: - result["name"] = self.name - - result["role"] = self.role.value - result["content"] = self.dump_content() - - if add_reasoning and self.reasoning_content: - result["reasoning_content"] = self.reasoning_content - - if self.tool_calls: - result["tool_calls"] = [tc.simple_output_dump() for tc in self.tool_calls] - - if self.tool_call_id: - result["tool_call_id"] = self.tool_call_id - - if add_time_created: - result["time_created"] = self.time_created - - if add_metadata: - result["metadata"] = self.metadata - - if enable_json_dump: - return json.dumps(result, ensure_ascii=False) - else: - return result - - def format_message( - self, - index: int | None = None, - add_time: bool = False, - use_name: bool = False, - add_reasoning: bool = True, - add_tools: bool = True, - strip_markdown_headers: bool = False, - ) -> str: - """Generates a human-readable string representation of the message.""" - prefix = f"round{index} " if index is not None else "" - time_str = f"[{self.time_created}] " if add_time else "" - header = f"{self.name or self.role.value if use_name else self.role.value}:" - - lines = [f"{prefix}{time_str}{header}"] - - def strip_md_func(line): - if strip_markdown_headers: - line = re.sub(r'\n##+ +', '\n', line) - return line - - if add_reasoning and self.reasoning_content: - lines.append(self.reasoning_content) - - if isinstance(self.content, str): - lines.append(strip_md_func(self.content)) - - elif isinstance(self.content, list): - for block in self.content: - text = block.content if isinstance(block.content, str) else \ - json.dumps(block.content, ensure_ascii=False) - text = str(text) - lines.append(strip_md_func(text)) - - if add_tools and self.tool_calls: - for tc in self.tool_calls: - lines.append(f" - tool_call={tc.name} params={tc.arguments}") - - return " ".join(lines).strip() - - -class Trajectory(BaseModel): - """Sequence of messages representing a full conversation session and its evaluation.""" - - task_id: str = Field(default="") - messages: list[Message] = Field(default_factory=list) - score: float = Field(default=0.0) - metadata: dict = Field(default_factory=dict) diff --git a/reme_ai/core/schema/request.py b/reme_ai/core/schema/request.py deleted file mode 100644 index ece942b9..00000000 --- a/reme_ai/core/schema/request.py +++ /dev/null @@ -1,11 +0,0 @@ -"""Defines the data structure for processing incoming user requests and message history.""" - -from pydantic import Field, BaseModel, ConfigDict - - -class Request(BaseModel): - """Represents a structured request payload containing a query, message list, and metadata.""" - - model_config = ConfigDict(extra="allow") - - metadata: dict = Field(default_factory=dict) diff --git a/reme_ai/core/schema/response.py b/reme_ai/core/schema/response.py deleted file mode 100644 index 3104bc6e..00000000 --- a/reme_ai/core/schema/response.py +++ /dev/null @@ -1,11 +0,0 @@ -"""Defines the standardized data structure for model output responses.""" - -from pydantic import Field, BaseModel - - -class Response(BaseModel): - """Represents a structured response containing the execution result, status, and metadata.""" - - answer: str | dict | list = Field(default="") - success: bool = Field(default=True) - metadata: dict = Field(default_factory=dict) diff --git a/reme_ai/core/schema/service_config.py b/reme_ai/core/schema/service_config.py deleted file mode 100644 index 4c6eb543..00000000 --- a/reme_ai/core/schema/service_config.py +++ /dev/null @@ -1,113 +0,0 @@ -"""Configuration schemas for service components using Pydantic models.""" - -import os -from typing import Dict, List - -from pydantic import BaseModel, Field, ConfigDict - -from .tool_call import ToolCall - - -class MCPConfig(BaseModel): - """Configuration for Model Context Protocol transport and network settings.""" - - model_config = ConfigDict(extra="allow") - - transport: str = Field(default="stdio") - host: str = Field(default="0.0.0.0") - port: int = Field(default=8001) - - -class HttpConfig(BaseModel): - """Configuration for the HTTP server interface and connection lifecycle.""" - - model_config = ConfigDict(extra="allow") - - host: str = Field(default="0.0.0.0") - port: int = Field(default=8001) - timeout_keep_alive: int = Field(default=3600) - limit_concurrency: int = Field(default=1000) - - -class CmdConfig(BaseModel): - """Configuration for command-line flow execution parameters.""" - - model_config = ConfigDict(extra="allow") - - flow: str = Field(default="") - - -class FlowConfig(ToolCall): - """Configuration for workflow execution, caching, and error handling.""" - - model_config = ConfigDict(extra="allow") - - flow_content: str = Field(default="") - stream: bool = Field(default=False) - raise_exception: bool = Field(default=True) - enable_cache: bool = Field(default=False) - cache_path: str = Field(default="cache/flow") - cache_expire_hours: float = Field(default=0.1) - - -class LLMConfig(BaseModel): - """Configuration for Large Language Model backend and model identification.""" - - model_config = ConfigDict(extra="allow") - - backend: str = Field(default="") - model_name: str = Field(default="") - - -class EmbeddingModelConfig(BaseModel): - """Configuration for embedding model backends and identity.""" - - model_config = ConfigDict(extra="allow") - - backend: str = Field(default="") - model_name: str = Field(default="") - - -class VectorStoreConfig(BaseModel): - """Configuration for vector database storage and associated embeddings.""" - - model_config = ConfigDict(extra="allow") - - backend: str = Field(default="local") - collection_name: str = Field(default="reme") - embedding_model: str = Field(default="default") - - -class TokenCounterConfig(BaseModel): - """Configuration for token counting services and model mapping.""" - - model_config = ConfigDict(extra="allow") - - backend: str = Field(default="base") - model_name: str = Field(default="") - - -class ServiceConfig(BaseModel): - """Root configuration schema aggregating all service-level settings and components.""" - - model_config = ConfigDict(extra="allow") - - backend: str = Field(default="") - app_name: str = Field(default=os.getenv("APP_NAME", "ReMe")) - enable_logo: bool = Field(default=True) - language: str = Field(default="") - thread_pool_max_workers: int = Field(default=16) - ray_max_workers: int = Field(default=-1) - init_logger: bool = Field(default=True) - disabled_flows: List[str] = Field(default_factory=list) - enabled_flows: List[str] = Field(default_factory=list) - mcp_servers: Dict[str, dict] = Field(default_factory=dict, description="External MCP Server configuration") - - mcp: MCPConfig = Field(default_factory=MCPConfig) - http: HttpConfig = Field(default_factory=HttpConfig) - cmd: CmdConfig = Field(default_factory=CmdConfig) - flow: Dict[str, FlowConfig] = Field(default_factory=dict) - llm: Dict[str, LLMConfig] = Field(default_factory=dict) - embedding_model: Dict[str, EmbeddingModelConfig] = Field(default_factory=dict) - vector_store: Dict[str, VectorStoreConfig] = Field(default_factory=dict) - token_counter: Dict[str, TokenCounterConfig] = Field(default_factory=dict) diff --git a/reme_ai/core/schema/stream_chunk.py b/reme_ai/core/schema/stream_chunk.py deleted file mode 100644 index 764981fd..00000000 --- a/reme_ai/core/schema/stream_chunk.py +++ /dev/null @@ -1,14 +0,0 @@ -"""Defines the data structure for individual data packets in a streaming response.""" - -from pydantic import Field, BaseModel - -from ..enumeration import ChunkEnum - - -class StreamChunk(BaseModel): - """Represents a single chunk of streamed data including its type, content, and completion status.""" - - chunk_type: ChunkEnum = Field(default=ChunkEnum.ANSWER) - chunk: str | dict | list = Field(default="") - done: bool = Field(default=False) - metadata: dict = Field(default_factory=dict) diff --git a/reme_ai/core/schema/tool_call.py b/reme_ai/core/schema/tool_call.py deleted file mode 100644 index 3c01cdcd..00000000 --- a/reme_ai/core/schema/tool_call.py +++ /dev/null @@ -1,228 +0,0 @@ -""" -MCP Tool Schema definitions for recursive JSON Schema representation. -""" - -import json -from typing import Any, Dict, List, Optional, Union - -from mcp.types import Tool -from pydantic import BaseModel, ConfigDict, Field, model_validator, field_validator - -from ..enumeration.json_schema_enum import JsonSchemaEnum - - -class ToolAttr(BaseModel): - """Recursive model representing JSON Schema attributes for tool parameters.""" - - model_config = ConfigDict(extra="allow") - - type: str = Field(default=str(JsonSchemaEnum.STRING), description="The data type of the attribute") - description: Optional[str] = Field(default=None, description="Description of the attribute") - required: Optional[List[str]] = Field(default=None, description="Required property names for object types") - properties: Optional[Dict[str, "ToolAttr"]] = Field(default=None, description="Child properties for objects") - items: Optional[Union[Dict[str, Any], "ToolAttr"]] = Field(default=None, description="Schema for array items") - enum: Optional[List[str]] = Field(default=None, description="Allowed values for the attribute") - - @field_validator("type") - @classmethod - def validate_type_is_valid_enum(cls, v: str) -> str: - """Validates that the provided type string exists within JsonSchemaEnum values.""" - valid_types = [str(e) for e in JsonSchemaEnum] - - if v not in valid_types: - raise ValueError(f"Invalid type: '{v}'. Must be one of {valid_types}") - return v - - def simple_input_dump(self) -> dict: - """Serializes the attribute into a standard JSON Schema dictionary.""" - res: dict = {"type": self.type} - if self.description: - res["description"] = self.description - if self.enum: - res["enum"] = self.enum - - if self.type == "object" and self.properties is not None: - res["properties"] = { - k: v.simple_input_dump() if isinstance(v, ToolAttr) else v for k, v in self.properties.items() - } - if self.required is not None: - res["required"] = self.required - - if self.type == "array" and self.items is not None: - res["items"] = self.items.simple_input_dump() if isinstance(self.items, ToolAttr) else self.items - - return res - - -# Enable recursive type resolution -ToolAttr.model_rebuild() - - -class ToolCall(BaseModel): - """ - Model representing a tool definition and its call structure. - Supports parsing from standard JSON Schema formats and converting to MCP Tool objects. - input: - { - "type": "function", - "function": { - "name": "get_current_weather", - "description": "It is very useful when you want to check the weather of a specified city.", - "parameters": { - "type": "object", - "properties": { - "location": { - "type": "string", - "description": "Cities or counties, such as Beijing, Hangzhou, Yuhang District, etc.", - } - }, - "required": ["location"] - } - } - } - output: - { - "index": 0, - "id": "call_6596dafa2a6a46f7a217da", - "function": { - "arguments": "{\"location\": \"Beijing\"}", - "name": "get_current_weather" - }, - "type": "function", - } - """ - - index: int = 0 - id: str = "" - type: str = "function" - name: str = "" - description: str = "" - - arguments: str = Field(default="", description="JSON string of tool execution arguments") - - parameters: ToolAttr = Field( - default_factory=lambda: ToolAttr(type="object", properties={}, required=[]), - description="Specification for input parameters", - ) - - output: ToolAttr = Field( - default_factory=lambda: ToolAttr(type="object", properties={}), - description="Specification for the execution result (Schema)", - ) - - @model_validator(mode="before") - @classmethod - def init_tool_call(cls, data: dict) -> dict: - """Initializes the model by parsing tool-specific body data.""" - data = data.copy() - t_type = data.get("type", "function") - body = data.get(t_type, {}) - - # Extract basic metadata - data["name"] = body.get("name", data.get("name", "")) - data["arguments"] = body.get("arguments", data.get("arguments", "")) - data["description"] = body.get("description", data.get("description", "")) - - # Handle parameters mapping - if "parameters" in body: - params = body["parameters"] - # If parameters is already a dict, ensure it matches ToolAttr structure - if isinstance(params, dict): - data["parameters"] = ToolAttr(**params) - - # Handle output mapping (if provided in source) - if "output" in body and isinstance(body["output"], dict): - data["output"] = ToolAttr(**body["output"]) - - return data - - def simple_input_dump(self) -> dict: - """Returns a standardized tool definition dictionary.""" - return { - "type": self.type, - self.type: { - "name": self.name, - "description": self.description, - "parameters": self.parameters.simple_input_dump(), - }, - } - - @classmethod - def from_mcp_tool(cls, tool: Tool) -> "ToolCall": - """Creates a ToolCall instance from an MCP Tool object.""" - # MCP Tool inputSchema maps directly to our parameters ToolAttr - return cls( - name=tool.name, - description=tool.description or "", - parameters=ToolAttr(**tool.inputSchema), - ) - - def to_mcp_tool(self) -> Tool: - """Converts the instance back into an MCP Tool object.""" - return Tool( - name=self.name, - description=self.description, - inputSchema=self.parameters.simple_input_dump(), - ) - - @property - def argument_dict(self) -> dict: - """Parse and return arguments as a dictionary.""" - return json.loads(self.arguments) - - def check_argument(self) -> bool: - """Check if arguments can be parsed as valid JSON.""" - try: - _ = self.argument_dict - return True - except Exception: - return False - - def sanitize_and_check_argument(self) -> bool: - """ - Attempt to sanitize and validate arguments JSON. - Common issues from LLM streaming: - - Extra closing brackets: }]}] -> }] - - Missing closing brackets - - Trailing commas - """ - if not self.arguments or not self.arguments.strip(): - return False - - try: - # First try parsing as-is - _ = json.loads(self.arguments) - return True - except json.JSONDecodeError: - pass - - # Try to fix common issues - sanitized = self.arguments.strip() - - # Remove trailing extra brackets/braces - # Pattern: if it ends with multiple closing chars, try removing extras - while len(sanitized) > 1: - try: - json.loads(sanitized) - self.arguments = sanitized # Update with sanitized version - return True - except json.JSONDecodeError: - # Try removing last character - if sanitized[-1] in ']}': - sanitized = sanitized[:-1].rstrip() - else: - break - - return False - - def simple_output_dump(self) -> dict: - """Convert ToolCall to output format dictionary for API responses.""" - return { - "index": self.index, - "id": self.id, - self.type: { - "arguments": self.arguments, - "name": self.name, - }, - "type": self.type, - } diff --git a/reme_ai/core/schema/vector_node.py b/reme_ai/core/schema/vector_node.py deleted file mode 100644 index 937ef4be..00000000 --- a/reme_ai/core/schema/vector_node.py +++ /dev/null @@ -1,15 +0,0 @@ -"""Defines the data structure for individual vector embedding nodes within a retrieval system.""" - -from typing import List, Dict -from uuid import uuid4 - -from pydantic import BaseModel, Field - - -class VectorNode(BaseModel): - """Represents a discrete unit of text content paired with its corresponding vector embedding and metadata.""" - - vector_id: str = Field(default_factory=lambda: uuid4().hex) - content: str = Field(default="") - vector: List[float] | None = Field(default=None) - metadata: Dict[str, str | bool | int | float] = Field(default_factory=dict) diff --git a/reme_ai/core/utils/__init__.py b/reme_ai/core/utils/__init__.py deleted file mode 100644 index 23396f97..00000000 --- a/reme_ai/core/utils/__init__.py +++ /dev/null @@ -1,47 +0,0 @@ -"""utils""" - -from .cache_handler import CacheHandler -from .case_converter import snake_to_camel, camel_to_snake -from .common_utils import run_coro_safely, execute_stream_task -from .env_utils import load_env -from .execute_tuils import exec_code, run_shell_command -from .http_client import HttpClient -from .llm_utils import extract_content, format_messages, deduplicate_memories -from .logger_utils import init_logger -from .logo_utils import print_logo - -# Make MCPClient import optional to avoid breaking if MCP dependencies are not available -try: - from .mcp_client import MCPClient - _HAS_MCP = True -except ImportError: - MCPClient = None - _HAS_MCP = False - -from .pydantic_config_parser import PydanticConfigParser -from .pydantic_utils import create_pydantic_model -from .singleton import singleton -from .time import timer, get_now_time - -__all__ = [ - "CacheHandler", - "snake_to_camel", - "camel_to_snake", - "run_coro_safely", - "execute_stream_task", - "load_env", - "exec_code", - "run_shell_command", - "HttpClient", - "extract_content", - "format_messages", - "deduplicate_memories", - "init_logger", - "print_logo", - "MCPClient", - "PydanticConfigParser", - "create_pydantic_model", - "singleton", - "timer", - "get_now_time", -] diff --git a/reme_ai/core/reme.py b/reme_ai/reme.py similarity index 99% rename from reme_ai/core/reme.py rename to reme_ai/reme.py index b5581159..6ec678c9 100644 --- a/reme_ai/core/reme.py +++ b/reme_ai/reme.py @@ -58,7 +58,6 @@ from .mem_tool.v4 import ( ) -@singleton class ReMe(Application): """Simplified ReMe application that auto-initializes the service context."""