mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-08 03:10:24 +00:00
feat(core): update version and enhance configuration management (#164)
This commit is contained in:
parent
8d7cc4bbd6
commit
d313ae6e52
6 changed files with 33 additions and 58 deletions
|
|
@ -6,7 +6,7 @@ from . import extension
|
|||
from . import memory
|
||||
from .reme import ReMe
|
||||
|
||||
__version__ = "0.3.0.8"
|
||||
__version__ = "0.3.0.9"
|
||||
|
||||
__all__ = [
|
||||
"config",
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@ from .registry_factory import R
|
|||
from .schema import Response, ServiceConfig
|
||||
from .service_context import ServiceContext
|
||||
from .token_counter import BaseTokenCounter
|
||||
from .utils import execute_stream_task, PydanticConfigParser, init_logger, MCPClient, print_logo, get_logger
|
||||
from .utils import execute_stream_task, PydanticConfigParser, init_logger, MCPClient, print_logo, get_logger, load_env
|
||||
from .vector_store import BaseVectorStore
|
||||
|
||||
logger = get_logger()
|
||||
|
|
@ -46,12 +46,16 @@ class Application:
|
|||
default_file_watcher_config: dict | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
|
||||
load_env()
|
||||
|
||||
self.llm_api_key = llm_api_key or os.getenv("LLM_API_KEY", "")
|
||||
self.llm_base_url = llm_base_url or os.getenv("LLM_BASE_URL", "")
|
||||
self.embedding_api_key = embedding_api_key or os.getenv("EMBEDDING_API_KEY", "")
|
||||
self.embedding_base_url = embedding_base_url or os.getenv("EMBEDDING_BASE_URL", "")
|
||||
|
||||
self.service_context = ServiceContext(
|
||||
*args,
|
||||
llm_api_key=llm_api_key,
|
||||
llm_base_url=llm_base_url,
|
||||
embedding_api_key=embedding_api_key,
|
||||
embedding_base_url=embedding_base_url,
|
||||
service_config=None,
|
||||
parser=parser,
|
||||
working_dir=working_dir,
|
||||
|
|
@ -158,11 +162,11 @@ class Application:
|
|||
else:
|
||||
config_dict = config.model_dump(exclude={"backend"})
|
||||
if not config_dict.get("api_key", ""):
|
||||
config_dict["api_key"] = os.getenv("LLM_API_KEY", "")
|
||||
config_dict["api_key"] = self.llm_api_key
|
||||
if "client_kwargs" not in config_dict:
|
||||
config_dict["client_kwargs"] = {}
|
||||
if not config_dict["client_kwargs"].get("base_url", ""):
|
||||
config_dict["client_kwargs"]["base_url"] = os.getenv("LLM_BASE_URL", "")
|
||||
config_dict["client_kwargs"]["base_url"] = self.llm_base_url
|
||||
self.service_context.as_llms[name] = R.as_llms[config.backend](**config_dict)
|
||||
|
||||
for name, config in self.service_config.as_llm_formatters.items():
|
||||
|
|
@ -184,6 +188,8 @@ class Application:
|
|||
logger.warning(f"LLM backend {config.backend} is not supported.")
|
||||
else:
|
||||
config_dict = config.model_dump(exclude={"backend"})
|
||||
config_dict.setdefault("api_key", self.llm_api_key)
|
||||
config_dict.setdefault("base_url", self.llm_base_url)
|
||||
self.service_context.llms[name] = R.llms[config.backend](**config_dict)
|
||||
await self.service_context.llms[name].start()
|
||||
|
||||
|
|
@ -192,7 +198,9 @@ class Application:
|
|||
logger.warning(f"Embedding model backend {config.backend} is not supported.")
|
||||
else:
|
||||
config_dict = config.model_dump(exclude={"backend"})
|
||||
config_dict["cache_dir"] = working_path / "embedding_cache"
|
||||
config_dict.setdefault("api_key", self.embedding_api_key)
|
||||
config_dict.setdefault("base_url", self.embedding_base_url)
|
||||
config_dict.setdefault("cache_dir", working_path / "embedding_cache")
|
||||
self.service_context.embedding_models[name] = R.embedding_models[config.backend](**config_dict)
|
||||
await self.service_context.embedding_models[name].start()
|
||||
|
||||
|
|
|
|||
|
|
@ -47,9 +47,18 @@ class ReMeTokenCounter(HuggingFaceTokenCounter):
|
|||
|
||||
# Set HuggingFace endpoint for mirror support
|
||||
if use_mirror:
|
||||
os.environ["HF_ENDPOINT"] = "https://hf-mirror.com"
|
||||
mirror = "https://hf-mirror.com"
|
||||
else:
|
||||
os.environ.pop("HF_ENDPOINT", None)
|
||||
mirror = "https://huggingface.co"
|
||||
|
||||
os.environ["HF_ENDPOINT"] = mirror
|
||||
|
||||
# if the huggingface is already imported in other dependencies,
|
||||
# we need to set the endpoint manually
|
||||
import huggingface_hub.constants
|
||||
|
||||
huggingface_hub.constants.ENDPOINT = mirror
|
||||
huggingface_hub.constants.HUGGINGFACE_CO_URL_TEMPLATE = mirror + "/{repo_id}/resolve/{revision}/{filename}"
|
||||
|
||||
try:
|
||||
super().__init__(
|
||||
|
|
|
|||
|
|
@ -6,7 +6,6 @@ Defines the abstract base class and standard API for all embedding model impleme
|
|||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from abc import ABC
|
||||
from collections import OrderedDict
|
||||
|
|
@ -56,8 +55,8 @@ class BaseEmbeddingModel(ABC):
|
|||
enable_cache: Whether to enable embedding cache
|
||||
**kwargs: Additional model-specific parameters
|
||||
"""
|
||||
self._api_key: str = api_key
|
||||
self._base_url: str = base_url
|
||||
self.api_key: str = api_key
|
||||
self.base_url: str = base_url
|
||||
self.model_name = model_name
|
||||
self.dimensions = dimensions
|
||||
self.use_dimensions = use_dimensions
|
||||
|
|
@ -78,16 +77,6 @@ class BaseEmbeddingModel(ABC):
|
|||
self.cache_path: Path = Path(self.cache_dir)
|
||||
self.cache_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
@property
|
||||
def api_key(self) -> str | None:
|
||||
"""Get API key from environment variable."""
|
||||
return os.getenv("EMBEDDING_API_KEY") or self._api_key
|
||||
|
||||
@property
|
||||
def base_url(self) -> str | None:
|
||||
"""Get base URL from environment variable."""
|
||||
return os.getenv("EMBEDDING_BASE_URL") or self._base_url
|
||||
|
||||
def _truncate_text(self, text: str) -> str:
|
||||
"""Truncate text to max_input_length if it exceeds the limit."""
|
||||
if len(text) > self.max_input_length:
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@
|
|||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Callable, Generator, AsyncGenerator, Any
|
||||
|
|
@ -36,8 +35,8 @@ class BaseLLM(ABC):
|
|||
request_interval: Minimum seconds between requests (default: 0.0)
|
||||
**kwargs: Additional model-specific parameters
|
||||
"""
|
||||
self._api_key: str = api_key
|
||||
self._base_url: str = base_url
|
||||
self.api_key: str = api_key
|
||||
self.base_url: str = base_url
|
||||
self.model_name: str = model_name
|
||||
self.max_retries: int = max_retries
|
||||
self.raise_exception: bool = raise_exception
|
||||
|
|
@ -47,16 +46,6 @@ class BaseLLM(ABC):
|
|||
self._last_request_time: float = 0.0
|
||||
self._request_lock: asyncio.Lock = asyncio.Lock()
|
||||
|
||||
@property
|
||||
def api_key(self) -> str | None:
|
||||
"""Get API key from environment variable."""
|
||||
return os.getenv("LLM_API_KEY") or self._api_key
|
||||
|
||||
@property
|
||||
def base_url(self) -> str | None:
|
||||
"""Get base URL from environment variable."""
|
||||
return os.getenv("LLM_BASE_URL") or self._base_url
|
||||
|
||||
@staticmethod
|
||||
def _accumulate_tool_call_chunk(tool_call, ret_tools: list[ToolCall]):
|
||||
"""Assemble incremental tool call chunks into complete ToolCall objects."""
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
"""Service context."""
|
||||
|
||||
import os
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
|
|
@ -8,7 +7,7 @@ from loguru import logger
|
|||
|
||||
from .base_dict import BaseDict
|
||||
from .schema import ServiceConfig
|
||||
from .utils import load_env, PydanticConfigParser
|
||||
from .utils import PydanticConfigParser
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agentscope.model import ChatModelBase
|
||||
|
|
@ -29,10 +28,6 @@ class ServiceContext(BaseDict):
|
|||
def __init__(
|
||||
self,
|
||||
*args,
|
||||
llm_api_key: str | None = None,
|
||||
llm_base_url: str | None = None,
|
||||
embedding_api_key: str | None = None,
|
||||
embedding_base_url: str | None = None,
|
||||
service_config: ServiceConfig | None = None,
|
||||
parser: type[PydanticConfigParser] | None = None,
|
||||
working_dir: str | None = None,
|
||||
|
|
@ -52,15 +47,6 @@ class ServiceContext(BaseDict):
|
|||
):
|
||||
super().__init__()
|
||||
|
||||
# Load environment variables
|
||||
load_env()
|
||||
|
||||
# Update common environment variables for LLM and embedding services.
|
||||
self.update_env("LLM_API_KEY", llm_api_key)
|
||||
self.update_env("LLM_BASE_URL", llm_base_url)
|
||||
self.update_env("EMBEDDING_API_KEY", embedding_api_key)
|
||||
self.update_env("EMBEDDING_BASE_URL", embedding_base_url)
|
||||
|
||||
if service_config is None:
|
||||
parser_class = parser if parser is not None else PydanticConfigParser
|
||||
parser_instance = parser_class(ServiceConfig)
|
||||
|
|
@ -114,12 +100,6 @@ class ServiceContext(BaseDict):
|
|||
self.flows: dict[str, "BaseFlow"] = {}
|
||||
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
|
||||
|
||||
@staticmethod
|
||||
def _update_section_config(config: dict, section_name: str, **kwargs):
|
||||
"""Update a specific section of the service config with new values."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue