mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-07 03:00:27 +00:00
feat(core): update API configuration handling with environment variable support
This commit is contained in:
parent
f9a5ad7b11
commit
35befd7bec
8 changed files with 64 additions and 39 deletions
|
|
@ -60,14 +60,18 @@ class Application:
|
|||
self.prompt_handler = PromptHandler(language=self.service_context.language)
|
||||
self._started: bool = False
|
||||
|
||||
def update_api_envs(self):
|
||||
"""Update the API environment variables."""
|
||||
self.service_context.update_api_envs()
|
||||
|
||||
@classmethod
|
||||
async def create(
|
||||
cls,
|
||||
*args,
|
||||
llm_api_key: str | None = None,
|
||||
llm_api_base: str | None = None,
|
||||
llm_base_url: str | None = None,
|
||||
embedding_api_key: str | None = None,
|
||||
embedding_api_base: str | None = None,
|
||||
embedding_base_url: str | None = None,
|
||||
enable_logo: bool = True,
|
||||
parser: type[PydanticConfigParser] | None = None,
|
||||
llm: dict | None = None,
|
||||
|
|
@ -82,9 +86,9 @@ class Application:
|
|||
instance = cls(
|
||||
*args,
|
||||
llm_api_key=llm_api_key,
|
||||
llm_base_url=llm_api_base,
|
||||
llm_base_url=llm_base_url,
|
||||
embedding_api_key=embedding_api_key,
|
||||
embedding_base_url=embedding_api_base,
|
||||
embedding_base_url=embedding_base_url,
|
||||
enable_logo=enable_logo,
|
||||
parser=parser,
|
||||
default_llm_config=llm,
|
||||
|
|
|
|||
|
|
@ -50,10 +50,7 @@ class ServiceContext(BaseContext):
|
|||
super().__init__()
|
||||
|
||||
load_env()
|
||||
self._update_env("REME_LLM_API_KEY", llm_api_key)
|
||||
self._update_env("REME_LLM_BASE_URL", llm_base_url)
|
||||
self._update_env("REME_EMBEDDING_API_KEY", embedding_api_key)
|
||||
self._update_env("REME_EMBEDDING_BASE_URL", embedding_base_url)
|
||||
self.update_api_envs(llm_api_key, llm_base_url, embedding_api_key, embedding_base_url)
|
||||
|
||||
if service_config is None:
|
||||
parser_class = parser if parser is not None else PydanticConfigParser
|
||||
|
|
@ -120,11 +117,24 @@ class ServiceContext(BaseContext):
|
|||
self._build_flows()
|
||||
|
||||
@staticmethod
|
||||
def _update_env(key: str, value: str | None):
|
||||
def update_env(key: str, value: str | None):
|
||||
"""Update environment variable if value is provided."""
|
||||
if value:
|
||||
os.environ[key] = value
|
||||
|
||||
def update_api_envs(
|
||||
self,
|
||||
llm_api_key: str | None = None,
|
||||
llm_base_url: str | None = None,
|
||||
embedding_api_key: str | None = None,
|
||||
embedding_base_url: str | None = None,
|
||||
):
|
||||
"""Update common environment variables for LLM and embedding services."""
|
||||
self.update_env("REME_LLM_API_KEY", llm_api_key)
|
||||
self.update_env("REME_LLM_BASE_URL", llm_base_url)
|
||||
self.update_env("REME_EMBEDDING_API_KEY", embedding_api_key)
|
||||
self.update_env("REME_EMBEDDING_BASE_URL", embedding_base_url)
|
||||
|
||||
@staticmethod
|
||||
def _update_section_config(config: dict, section_name: str, **kwargs):
|
||||
"""Update a specific section of the service config with new values."""
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ Defines the abstract base class and standard API for all embedding model impleme
|
|||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import os
|
||||
import time
|
||||
from abc import ABC
|
||||
from collections import OrderedDict
|
||||
|
|
@ -24,6 +25,8 @@ class BaseEmbeddingModel(ABC):
|
|||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str | None = None,
|
||||
base_url: str | None = None,
|
||||
model_name: str = "",
|
||||
dimensions: int | None = 1024,
|
||||
max_batch_size: int = 10,
|
||||
|
|
@ -36,6 +39,8 @@ class BaseEmbeddingModel(ABC):
|
|||
"""Initialize model configuration and parameters.
|
||||
|
||||
Args:
|
||||
api_key: API key for the embedding service
|
||||
base_url: Base URL for the embedding service
|
||||
model_name: Name of the embedding model
|
||||
dimensions: Vector dimensions of the embeddings
|
||||
max_batch_size: Maximum batch size for embedding requests
|
||||
|
|
@ -45,6 +50,8 @@ class BaseEmbeddingModel(ABC):
|
|||
max_cache_size: Maximum number of embeddings to cache in memory (LRU)
|
||||
**kwargs: Additional model-specific parameters
|
||||
"""
|
||||
self._api_key: str = api_key
|
||||
self._base_url: str = base_url
|
||||
self.model_name = model_name
|
||||
self.dimensions = dimensions
|
||||
self.max_batch_size = max_batch_size
|
||||
|
|
@ -59,6 +66,16 @@ class BaseEmbeddingModel(ABC):
|
|||
self._cache_hits = 0
|
||||
self._cache_misses = 0
|
||||
|
||||
@property
|
||||
def api_key(self) -> str | None:
|
||||
"""Get API key from environment variable."""
|
||||
return os.getenv("REME_EMBEDDING_API_KEY") or self._api_key
|
||||
|
||||
@property
|
||||
def base_url(self) -> str | None:
|
||||
"""Get base URL from environment variable."""
|
||||
return os.getenv("REME_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:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
"""Asynchronous OpenAI-compatible embedding model implementation for ReMe."""
|
||||
|
||||
import os
|
||||
from typing import Literal
|
||||
|
||||
from openai import AsyncOpenAI
|
||||
|
|
@ -11,17 +10,9 @@ from .base_embedding_model import BaseEmbeddingModel
|
|||
class OpenAIEmbeddingModel(BaseEmbeddingModel):
|
||||
"""Asynchronous embedding model implementation compatible with OpenAI-style APIs."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str | None = None,
|
||||
base_url: str | None = None,
|
||||
encoding_format: Literal["float", "base64"] = "float",
|
||||
**kwargs,
|
||||
):
|
||||
def __init__(self, encoding_format: Literal["float", "base64"] = "float", **kwargs):
|
||||
"""Initialize the OpenAI async embedding model with API credentials and configuration."""
|
||||
super().__init__(**kwargs)
|
||||
self.api_key: str = api_key or os.getenv("REME_EMBEDDING_API_KEY", "")
|
||||
self.base_url: str = base_url or os.getenv("REME_EMBEDDING_BASE_URL", "")
|
||||
self.encoding_format: Literal["float", "base64"] = encoding_format
|
||||
|
||||
# Create client using factory method
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@
|
|||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Callable, Generator, AsyncGenerator, Any
|
||||
|
|
@ -20,7 +21,9 @@ class BaseLLM(ABC):
|
|||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
api_key: str | None = None,
|
||||
base_url: str | None = None,
|
||||
model_name: str = "",
|
||||
max_retries: int = 10,
|
||||
raise_exception: bool = False,
|
||||
request_interval: float = 0.0,
|
||||
|
|
@ -35,6 +38,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.model_name: str = model_name
|
||||
self.max_retries: int = max_retries
|
||||
self.raise_exception: bool = raise_exception
|
||||
|
|
@ -44,6 +49,16 @@ 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("REME_LLM_API_KEY") or self._api_key
|
||||
|
||||
@property
|
||||
def base_url(self) -> str | None:
|
||||
"""Get base URL from environment variable."""
|
||||
return os.getenv("REME_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 @@
|
|||
"""LiteLLM asynchronous implementation for ReMe."""
|
||||
|
||||
import os
|
||||
from typing import AsyncGenerator
|
||||
|
||||
from loguru import logger
|
||||
|
|
@ -15,17 +14,9 @@ from ..schema import ToolCall
|
|||
class LiteLLM(BaseLLM):
|
||||
"""Async LLM implementation using LiteLLM to support multiple providers."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str | None = None,
|
||||
base_url: str | None = None,
|
||||
custom_llm_provider: str = "openai",
|
||||
**kwargs,
|
||||
):
|
||||
def __init__(self, custom_llm_provider: str = "openai", **kwargs):
|
||||
"""Initialize the LiteLLM client with API configuration and provider settings."""
|
||||
super().__init__(**kwargs)
|
||||
self.api_key: str | None = api_key or os.getenv("REME_LLM_API_KEY")
|
||||
self.base_url: str | None = base_url or os.getenv("REME_LLM_BASE_URL")
|
||||
self.custom_llm_provider: str = custom_llm_provider
|
||||
|
||||
def _build_stream_kwargs(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
"""Asynchronous OpenAI-compatible LLM implementation supporting streaming, tool calls, and reasoning content."""
|
||||
|
||||
import os
|
||||
from typing import AsyncGenerator
|
||||
|
||||
from loguru import logger
|
||||
|
|
@ -16,16 +15,9 @@ from ..schema import ToolCall
|
|||
class OpenAILLM(BaseLLM):
|
||||
"""Asynchronous LLM client for OpenAI-compatible APIs supporting streaming completions and tool execution."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str | None = None,
|
||||
base_url: str | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
def __init__(self, **kwargs):
|
||||
"""Initialize the OpenAI async client with API credentials and model configuration."""
|
||||
super().__init__(**kwargs)
|
||||
self.api_key: str = api_key or os.getenv("REME_LLM_API_KEY", "")
|
||||
self.base_url: str = base_url or os.getenv("REME_LLM_BASE_URL", "")
|
||||
|
||||
# Create client using factory method
|
||||
self._client = self._create_client()
|
||||
|
|
|
|||
|
|
@ -70,6 +70,11 @@ class SqliteMemoryStore(BaseMemoryStore):
|
|||
|
||||
# Only load sqlite-vec extension if vector search is enabled
|
||||
if self.vector_enabled:
|
||||
logger.warning(
|
||||
"On macOS systems with version 14 or earlier, "
|
||||
"loading the sqlite-vec vector extension carries a risk of crashes or hangs.",
|
||||
)
|
||||
|
||||
self.conn.enable_load_extension(True)
|
||||
|
||||
# Load sqlite-vec extension
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue