feat(core): update API configuration handling with environment variable support

This commit is contained in:
jinli.yl 2026-02-12 16:11:09 +08:00
parent f9a5ad7b11
commit 35befd7bec
8 changed files with 64 additions and 39 deletions

View file

@ -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,

View file

@ -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."""

View file

@ -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:

View file

@ -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

View file

@ -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."""

View file

@ -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(

View file

@ -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()

View file

@ -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