ReMe/reme/core/service_context.py
Zhouwk 625d184ca1
Some checks failed
Pre-commit / run (ubuntu-latest) (push) Has been cancelled
添加log_to_file的开关 (#205)
2026-04-13 10:49:26 +08:00

112 lines
4.8 KiB
Python

"""Service context."""
from concurrent.futures import ThreadPoolExecutor
from typing import TYPE_CHECKING
from loguru import logger
from .base_dict import BaseDict
from .schema import ServiceConfig
from .utils import PydanticConfigParser
if TYPE_CHECKING:
from agentscope.model import ChatModelBase
from agentscope.formatter import FormatterBase
from agentscope.token import TokenCounterBase
from .llm import BaseLLM
from .embedding import BaseEmbeddingModel
from .vector_store import BaseVectorStore
from .file_store import BaseFileStore
from .token_counter import BaseTokenCounter
from .flow import BaseFlow
from .file_watcher import BaseFileWatcher
class ServiceContext(BaseDict):
"""Service context."""
def __init__(
self,
*args,
service_config: ServiceConfig | None = None,
parser: type[PydanticConfigParser] | None = None,
working_dir: str | None = None,
config_path: str | None = None,
enable_logo: bool = True,
log_to_console: bool = True,
log_to_file: bool = True,
default_as_llm_config: dict | None = None,
default_as_llm_formatter_config: dict | None = None,
default_as_token_counter_config: dict | None = None,
default_llm_config: dict | None = None,
default_embedding_model_config: dict | None = None,
default_vector_store_config: dict | None = None,
default_file_store_config: dict | None = None,
default_token_counter_config: dict | None = None,
default_file_watcher_config: dict | None = None,
**kwargs,
):
super().__init__()
if service_config is None:
parser_class = parser if parser is not None else PydanticConfigParser
parser_instance = parser_class(ServiceConfig)
input_args = []
if config_path:
input_args.append(f"config={config_path}")
if args:
input_args.extend(args)
if default_as_llm_config:
self._update_section_config(kwargs, "as_llms", **default_as_llm_config)
if default_as_llm_formatter_config:
self._update_section_config(kwargs, "as_llm_formatters", **default_as_llm_formatter_config)
if default_as_token_counter_config:
self._update_section_config(kwargs, "as_token_counters", **default_as_token_counter_config)
if default_llm_config:
self._update_section_config(kwargs, "llms", **default_llm_config)
if default_embedding_model_config:
self._update_section_config(kwargs, "embedding_models", **default_embedding_model_config)
if default_token_counter_config:
self._update_section_config(kwargs, "token_counters", **default_token_counter_config)
if default_vector_store_config:
self._update_section_config(kwargs, "vector_stores", **default_vector_store_config)
if default_file_store_config:
self._update_section_config(kwargs, "file_stores", **default_file_store_config)
if default_file_watcher_config:
self._update_section_config(kwargs, "file_watchers", **default_file_watcher_config)
kwargs.update(
{
"enable_logo": enable_logo,
"log_to_console": log_to_console,
"log_to_file": log_to_file,
"working_dir": working_dir,
},
)
logger.info(f"update with args: {input_args} kwargs: {kwargs}")
service_config = parser_instance.parse_args(*input_args, **kwargs)
self.service_config: ServiceConfig = service_config
self.thread_pool: ThreadPoolExecutor | None = None
self.as_llms: dict[str, "ChatModelBase"] = {}
self.as_llm_formatters: dict[str, "FormatterBase"] = {}
self.as_token_counters: dict[str, "TokenCounterBase"] = {}
self.llms: dict[str, "BaseLLM"] = {}
self.embedding_models: dict[str, "BaseEmbeddingModel"] = {}
self.token_counters: dict[str, "BaseTokenCounter"] = {}
self.vector_stores: dict[str, "BaseVectorStore"] = {}
self.file_stores: dict[str, "BaseFileStore"] = {}
self.file_watchers: dict[str, "BaseFileWatcher"] = {}
self.flows: dict[str, "BaseFlow"] = {}
self.mcp_server_mapping: dict[str, dict] = {}
@staticmethod
def _update_section_config(config: dict, section_name: str, **kwargs):
"""Update a specific section of the service config with new values."""
if section_name not in config:
config[section_name] = {}
if "default" not in config[section_name]:
config[section_name]["default"] = {}
config[section_name]["default"].update(kwargs)