feat(llm): add rate limiting and model override support to BaseLLM

This commit is contained in:
jinli.yl 2026-01-11 11:24:20 +08:00
parent 2e2d96a53e
commit e302ec72cb
8 changed files with 394 additions and 36 deletions

View file

@ -9,7 +9,9 @@ This script performs the full evaluation pipeline:
Usage:
python bench/halumem/eval_reme.py --data_path /Users/yuli/workspace/HaluMem/data/HaluMem-Long.jsonl --version v1 \
--top_k 20 --user_num 10 --max_concurrency 5
--top_k 20 --user_num 100 --max_concurrency 20
python bench/halumem/eval_reme.py --data_path ./HaluMem-Long.jsonl --version v1 \
--top_k 20 --user_num 100 --max_concurrency 20
"""
import asyncio
@ -40,6 +42,8 @@ TEMPLATE_MEMOS = _PROMPTS["TEMPLATE_MEMOS"]
# Prompt for question answering (using optimized PROMPT_MEMOS)
PROMPT_MEMOS = _PROMPTS["PROMPT_MEMOS"]
# Initialize ReMe with rate limiting configuration
# The default LLM can be overridden at call time using model_name parameter
reme: ReMe = ReMe()
@ -235,12 +239,15 @@ async def process_user_stage1(
if extract_memories_str.strip() == "":
new_memory["memory_integrity_score"] = 0
new_memory["memory_integrity_reasoning"] = "No memories extracted"
session_eval_results["memory_integrity_records"].append(new_memory)
continue
result = await evaluation_for_memory_integrity(extract_memories_str, memory["memory_content"])
score = int(result.get("score"))
reasoning = result.get("reasoning", "")
new_memory["memory_integrity_score"] = score
new_memory["memory_integrity_reasoning"] = reasoning
session_eval_results["memory_integrity_records"].append(new_memory)
# Evaluate Memory Accuracy
@ -266,8 +273,10 @@ async def process_user_stage1(
result = await evaluation_for_memory_accuracy(dialogue_str, golden_memories_str, memory)
score = int(result.get("accuracy_score"))
is_included_in_golden_memories = result.get("is_included_in_golden_memories", "false")
reason = result.get("reason", "")
new_memory["memory_accuracy_score"] = score
new_memory["is_included_in_golden_memories"] = is_included_in_golden_memories
new_memory["memory_accuracy_reason"] = reason
session_eval_results["memory_accuracy_records"].append(new_memory)
# Evaluate Memory Update
@ -289,7 +298,9 @@ async def process_user_stage1(
"\n".join(update_memory["original_memories"]),
)
update_type = result.get("evaluation_result")
reason = result.get("reason", "")
update_memory["memory_update_type"] = update_type
update_memory["memory_update_reason"] = reason
session_eval_results["memory_update_records"].append(update_memory)
# Evaluate Question Answering
@ -307,7 +318,9 @@ async def process_user_stage1(
qa["system_response"],
)
result_type = result.get("evaluation_result")
reasoning = result.get("reasoning", "")
new_qa["result_type"] = result_type
new_qa["question_answering_reasoning"] = reasoning
session_eval_results["question_answering_records"].append(new_qa)
# Store evaluation results in session
@ -676,13 +689,7 @@ async def main_async(
semaphore_stage1 = asyncio.Semaphore(max_concurrency)
async def process_single_user_stage1(idx: int, user_data: dict):
"""Process a single user in Stage 1 with semaphore control and staggered delay."""
# Add staggered delay: 0s for first, 30s for second, 60s for third, etc.
delay = (idx - 1) * 30
if delay > 0:
print(f"⏳ User {idx} will start in {delay} seconds...")
await asyncio.sleep(delay)
"""Process a single user in Stage 1 with semaphore control."""
async with semaphore_stage1:
uuid = user_data['uuid']
tmp_file = os.path.join(tmp_dir, f"{uuid}.json")
@ -696,7 +703,7 @@ async def main_async(
print(f"[{idx}/{total_users}] ✅ Finished {uuid} ({result['status']})")
return result
# Process users in parallel with controlled concurrency and staggered start
# Process users in parallel with controlled concurrency
tasks = [process_single_user_stage1(idx, user_data) for idx, user_data in enumerate(user_data_list, 1)]
await asyncio.gather(*tasks)

View file

@ -418,4 +418,77 @@ EVALUATION_PROMPT_FOR_QUESTION: |
"reasoning": "Provide a concise and traceable evaluation rationale: first compare the systems response with the Key Memory Points (which were correctly used, which were missing, and whether there was any fabrication/contradiction), then assess its consistency with the Reference Answer, and finally state the classification basis.",
"evaluation_result": "Correct | Hallucination | Omission"
}}
```
```
EVALUATION_PROMPT_FOR_QUESTION2: |
You are an **evaluation expert for AI memory system question answering**.
* **Dialogue:**
{dialogue}
Based **only** on the provided **“Question”**, **“Reference Answer”**, and **“Key Memory Points”** (the essential facts needed to derive the reference answer), strictly evaluate the **accuracy** of the **“Memory System Response.”** Classify it as one of **“Correct”**, **“Hallucination”**, or **“Omission.”** Do **not** use any external knowledge or subjective inference. Finally, output your judgment **strictly** in the specified JSON format.
# Evaluation Criteria
## Answer Type Classification
### 1. Correct
* The “Memory System Response” accurately answers the “Question,” and its content is **semantically equivalent** to the “Reference Answer.”
* It contains **no contradictions** with the “Key Memory Points” or “Reference Answer.”
* It introduces **no unsupported details** beyond the “Key Memory Points” that could alter the conclusion.
* Synonyms, paraphrasing, and reasonable summarization are acceptable.
### 2. Hallucination
* The “Memory System Response” includes information or facts that **contradict or are inconsistent** with the “Reference Answer” or the “Key Memory Points.”
* When the “Reference Answer” is labeled as *unknown/uncertain*, yet the response provides a specific verifiable fact or conclusion.
* Extra irrelevant information that does **not change** the conclusion is **not** considered hallucination by itself; however, if it **changes or misleads** the conclusion, or **contradicts** the “Key Memory Points,” it should be judged as a **Hallucination**.
### 3. Omission
* The response is **incomplete** compared to the “Reference Answer.”
* It explicitly states “dont know,” “cant remember,” or “no related memory,” even though relevant information exists in the “Key Memory Points.”
* For multi-element questions, **all elements must be correct and present**; omission of **any** element is considered an **Omission**.
## Priority Rules (Conflict Handling)
* If the response contains **both missing necessary information** and **fabricated/contradictory information**, classify it as **Hallucination**.
* If there is **no fabrication/contradiction** but some necessary information is missing, classify it as **Omission**.
* Only when the meaning is **fully equivalent** to the reference answer should it be classified as **Correct**.
## Detailed Guidelines and Tolerance
* Equivalent expressions of numbers, times, and units are acceptable, but the **numerical values themselves must not differ**.
* For multi-element questions, **all elements must be complete and accurate**; missing any element counts as **Omission**.
* If the reference answer is *“unknown / cannot be determined”* and the system provides a definite fact, that is a **Hallucination**.
If the system also answers *“unknown”* (without guessing), it may be **Correct**.
* The evaluation must rely **only** on the *Reference Answer*, *Key Memory Points*, and *System Response* — no external context, world knowledge, or speculative reasoning is allowed.
# Information for Evaluation
* **Question:**
{question}
* **Reference Answer:**
{reference_answer}
* **Key Memory Points:**
{key_memory_points}
* **Memory System Response:**
{response}
# Output Requirements
Please provide your evaluation result **strictly** in the JSON format below.
Do **not** add any extra explanation or comments outside the JSON block.
```json
{{
"reasoning": "Provide a concise and traceable evaluation rationale: first compare the systems response with the Key Memory Points (which were correctly used, which were missing, and whether there was any fabrication/contradiction), then assess its consistency with the Reference Answer, and finally state the classification basis.",
"evaluation_result": "Correct | Hallucination | Omission"
}}
```
"""

View file

@ -5,9 +5,9 @@ import re
from tenacity import retry, stop_after_attempt, wait_random_exponential, before_sleep_log
from reme_ai.core.llm import OpenAILLM
from reme_ai.core.schema import Message
from reme_ai.core.utils import load_env
from reme_ai.reme import ReMe
logger = logging.getLogger(__name__)
@ -17,6 +17,8 @@ WAIT_TIME_LOWER = 1
WAIT_TIME_UPPER = 60
RETRY_TIMES = 5
# Use ReMe singleton's LLM instead of creating a separate instance
reme = ReMe()
@retry(
wait=wait_random_exponential(min=WAIT_TIME_LOWER, max=WAIT_TIME_UPPER),
@ -24,9 +26,18 @@ RETRY_TIMES = 5
reraise=True,
before_sleep=before_sleep_log(logger, logging.WARNING),
)
async def llm_request(prompt, **kwargs) -> str:
llm = OpenAILLM(model_name="qwen3-max")
assistant_message = await llm.chat(
async def llm_request(prompt, model_name: str = "qwen3-max", **kwargs) -> str:
"""Make an LLM request using ReMe's LLM with optional model override.
Args:
prompt: The prompt to send to the LLM
model_name: Optional model name to override the default model (default: "qwen3-max")
**kwargs: Additional arguments to pass to the chat method
Returns:
The assistant's response content
"""
assistant_message = await reme.llm.chat(
messages=[
Message(
**{
@ -35,6 +46,7 @@ async def llm_request(prompt, **kwargs) -> str:
},
),
],
model_name=model_name,
**kwargs,
)
return assistant_message.content
@ -46,8 +58,21 @@ async def llm_request(prompt, **kwargs) -> str:
reraise=True,
before_sleep=before_sleep_log(logger, logging.WARNING),
)
async def llm_request_for_json(prompt, **kwargs):
content = await llm_request(prompt, **kwargs)
async def llm_request_for_json(prompt, model_name: str = "qwen3-max", **kwargs):
"""Make an LLM request expecting JSON response using ReMe's LLM.
Args:
prompt: The prompt to send to the LLM
model_name: Optional model name to override the default model (default: "qwen3-max")
**kwargs: Additional arguments to pass to the chat method
Returns:
Parsed JSON object from the LLM response
Raises:
ValueError: If no JSON block is found in the model output
"""
content = await llm_request(prompt, model_name=model_name, **kwargs)
match = re.search(r"```json\s*(\{.*?\})\s*```", content, re.DOTALL)
if not match:

View file

@ -16,11 +16,15 @@ llm:
default:
backend: openai
model_name: qwen3-30b-a3b-instruct-2507
max_rps: 6
rps_window: 10
qwen3_max_instruct:
backend: openai
model_name: qwen3-max
temperature: 0.6
# temperature: 0.6
max_rps: 9
rps_window: 10
embedding_model:
default:
@ -31,6 +35,7 @@ embedding_model:
vector_store:
default:
backend: chroma
# backend: local
embedding_model: default
collection_name: reme

View file

@ -4,6 +4,7 @@ import asyncio
import json
import time
from abc import ABC, abstractmethod
from collections import deque
from typing import Callable, Generator, AsyncGenerator, Any
from loguru import logger
@ -17,12 +18,90 @@ from ..schema import ToolCall
class BaseLLM(ABC):
"""Abstract base class defining the standard interface for LLM interactions."""
def __init__(self, model_name: str, max_retries: int = 3, raise_exception: bool = False, **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, max_rps: int | None = None, rps_window: float = 1.0, **kwargs):
"""Initialize the LLM client with model configurations and retry policies.
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
max_rps: Maximum requests allowed within the time window. If None, no rate limiting is applied.
rps_window: Time window in seconds for rate limiting (default: 1.0).
For example: max_rps=10, rps_window=5.0 means max 10 requests in 5 seconds.
**kwargs: Additional model-specific parameters
"""
self.model_name: str = model_name
self.max_retries: int = max_retries
self.raise_exception: bool = raise_exception
self.max_rps: int | None = max_rps
self.rps_window: float = rps_window
self.kwargs: dict = kwargs
# Rate limiting state - using deque for efficient O(1) operations
self._request_timestamps: deque = deque()
self._rate_limit_lock = asyncio.Lock() # For async rate limiting
import threading
self._rate_limit_lock_sync = threading.Lock() # For sync rate limiting
async def _wait_for_rate_limit(self):
"""Async rate limiting: wait if necessary to respect max_rps constraint within the time window."""
if self.max_rps is None:
return
async with self._rate_limit_lock:
current_time = time.time()
# Remove timestamps older than the time window
while self._request_timestamps and current_time - self._request_timestamps[0] >= self.rps_window:
self._request_timestamps.popleft()
# If we've reached the rate limit, wait until we can proceed
if len(self._request_timestamps) >= self.max_rps:
# Calculate how long to wait
oldest_timestamp = self._request_timestamps[0]
wait_time = self.rps_window - (current_time - oldest_timestamp)
if wait_time > 0:
logger.debug(f"Rate limit reached ({self.max_rps} requests in {self.rps_window}s). Waiting {wait_time:.3f}s")
await asyncio.sleep(wait_time)
# Clean up old timestamps after waiting
current_time = time.time()
while self._request_timestamps and current_time - self._request_timestamps[0] >= self.rps_window:
self._request_timestamps.popleft()
# Record this request
self._request_timestamps.append(time.time())
def _wait_for_rate_limit_sync(self):
"""Synchronous rate limiting: wait if necessary to respect max_rps constraint within the time window."""
if self.max_rps is None:
return
with self._rate_limit_lock_sync:
current_time = time.time()
# Remove timestamps older than the time window
while self._request_timestamps and current_time - self._request_timestamps[0] >= self.rps_window:
self._request_timestamps.popleft()
# If we've reached the rate limit, wait until we can proceed
if len(self._request_timestamps) >= self.max_rps:
# Calculate how long to wait
oldest_timestamp = self._request_timestamps[0]
wait_time = self.rps_window - (current_time - oldest_timestamp)
if wait_time > 0:
logger.debug(f"Rate limit reached ({self.max_rps} requests in {self.rps_window}s). Waiting {wait_time:.3f}s")
time.sleep(wait_time)
# Clean up old timestamps after waiting
current_time = time.time()
while self._request_timestamps and current_time - self._request_timestamps[0] >= self.rps_window:
self._request_timestamps.popleft()
# Record this request
self._request_timestamps.append(time.time())
@staticmethod
def _accumulate_tool_call_chunk(tool_call, ret_tools: list[ToolCall]):
@ -68,9 +147,18 @@ class BaseLLM(ABC):
messages: list[Message],
tools: list[ToolCall] | None = None,
log_params: bool = True,
model_name: str | None = None,
**kwargs,
) -> dict:
"""Construct provider-specific parameters for streaming API requests."""
"""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
"""
async def _stream_chat(
self,
@ -94,10 +182,21 @@ class BaseLLM(ABC):
self,
messages: list[Message],
tools: list[ToolCall] | None = None,
model_name: str | None = None,
**kwargs,
) -> AsyncGenerator[StreamChunk, None]:
"""Public async interface for streaming chat completions with retries."""
stream_kwargs = self._build_stream_kwargs(messages, tools, **kwargs)
"""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 rate limiting before making the request
await self._wait_for_rate_limit()
stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs)
for i in range(self.max_retries):
try:
@ -121,10 +220,21 @@ class BaseLLM(ABC):
self,
messages: list[Message],
tools: list[ToolCall] | None = None,
model_name: str | None = None,
**kwargs,
) -> Generator[StreamChunk, None, None]:
"""Public synchronous interface for streaming chat completions with retries."""
stream_kwargs = self._build_stream_kwargs(messages, tools, **kwargs)
"""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
"""
# Apply rate limiting before making the request
self._wait_for_rate_limit_sync()
stream_kwargs = self._build_stream_kwargs(messages, tools, model_name=model_name, **kwargs)
for i in range(self.max_retries):
try:
@ -148,9 +258,18 @@ class BaseLLM(ABC):
messages: list[Message],
tools: list[ToolCall] | None = None,
enable_stream_print: bool = False,
model_name: str | None = None,
**kwargs,
) -> Message:
"""Internal async method to aggregate a full response by consuming the stream."""
"""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
"""
state = {
"enter_think": False,
"enter_answer": False,
@ -159,7 +278,7 @@ class BaseLLM(ABC):
"tool_calls": [],
}
stream_kwargs = self._build_stream_kwargs(messages, tools, **kwargs)
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:
@ -207,9 +326,18 @@ class BaseLLM(ABC):
messages: list[Message],
tools: list[ToolCall] | None = None,
enable_stream_print: bool = False,
model_name: str | None = None,
**kwargs,
) -> Message:
"""Internal synchronous method to aggregate a full response by consuming the stream."""
"""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
"""
state = {
"enter_think": False,
"enter_answer": False,
@ -218,7 +346,7 @@ class BaseLLM(ABC):
"tool_calls": [],
}
stream_kwargs = self._build_stream_kwargs(messages, tools, **kwargs)
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:
@ -268,21 +396,66 @@ class BaseLLM(ABC):
enable_stream_print: bool = False,
callback_fn: Callable[[Message], Any] | None = None,
default_value: Any = None,
model_name: str | None = None,
**kwargs,
) -> Message | Any:
"""Perform an async chat completion with integrated retries and error handling."""
"""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
"""
# Use the provided model_name or fall back to self.model_name
effective_model = model_name if model_name is not None else self.model_name
for i in range(self.max_retries):
try:
# Apply rate limiting before making the request
await self._wait_for_rate_limit()
result = await self._chat(
messages=messages,
tools=tools,
enable_stream_print=enable_stream_print,
model_name=model_name,
**kwargs,
)
return callback_fn(result) if callback_fn else result
except Exception as e:
logger.exception(f"chat with model={self.model_name} encounter error with e={e.args}")
# 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()
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("=" * 80)
for idx, msg in enumerate(messages):
logger.error(f"Message {idx + 1} [role={msg.role}]:")
logger.error(f"Content: {msg.content}")
if msg.reasoning_content:
logger.error(f"Reasoning: {msg.reasoning_content}")
if msg.tool_calls:
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})")
await asyncio.sleep(60)
continue
logger.exception(f"chat with model={effective_model} encounter error with e={e.args}")
if i == self.max_retries - 1:
if self.raise_exception:
@ -299,21 +472,66 @@ class BaseLLM(ABC):
enable_stream_print: bool = False,
callback_fn: Callable[[Message], Any] | None = None,
default_value: Any = None,
model_name: str | None = None,
**kwargs,
) -> Message | Any:
"""Perform a synchronous chat completion with integrated retries and error handling."""
"""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
effective_model = model_name if model_name is not None else self.model_name
for i in range(self.max_retries):
try:
# Apply rate limiting before making the request
self._wait_for_rate_limit_sync()
result = self._chat_sync(
messages=messages,
tools=tools,
enable_stream_print=enable_stream_print,
model_name=model_name,
**kwargs,
)
return callback_fn(result) if callback_fn else result
except Exception as e:
logger.exception(f"chat sync with model={self.model_name} encounter error with e={e.args}")
# 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()
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("=" * 80)
for idx, msg in enumerate(messages):
logger.error(f"Message {idx + 1} [role={msg.role}]:")
logger.error(f"Content: {msg.content}")
if msg.reasoning_content:
logger.error(f"Reasoning: {msg.reasoning_content}")
if msg.tool_calls:
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})")
time.sleep(60)
continue
logger.exception(f"chat sync with model={effective_model} encounter error with e={e.args}")
if i == self.max_retries - 1:
if self.raise_exception:

View file

@ -36,12 +36,24 @@ class LiteLLM(BaseLLM):
messages: list[Message],
tools: list[ToolCall] | None = None,
log_params: bool = True,
model_name: str | None = None,
**kwargs,
) -> dict:
"""Construct and log the parameters dictionary for LiteLLM API calls."""
"""Construct and log the parameters dictionary for LiteLLM API calls.
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
"""
# Use the provided model_name or fall back to self.model_name
effective_model = model_name if model_name is not None else self.model_name
# Construct the API parameters by merging multiple sources
llm_kwargs = {
"model": self.model_name,
"model": effective_model,
"messages": [x.simple_dump() for x in messages],
"tools": [x.simple_input_dump() for x in tools] if tools else None,
"stream": True,

View file

@ -41,12 +41,24 @@ class OpenAILLM(BaseLLM):
messages: list[Message],
tools: list[ToolCall] | None = None,
log_params: bool = True,
model_name: str | None = None,
**kwargs,
) -> dict:
"""Construct the parameter dictionary for the OpenAI Chat Completions API call."""
"""Construct the parameter dictionary for the OpenAI Chat Completions API call.
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
"""
# Use the provided model_name or fall back to self.model_name
effective_model = model_name if model_name is not None else self.model_name
# Construct the API parameters by merging multiple sources
llm_kwargs = {
"model": self.model_name,
"model": effective_model,
"messages": [x.simple_dump() for x in messages],
"tools": [x.simple_input_dump() for x in tools] if tools else None,
"stream": True,

View file

@ -7,6 +7,7 @@ from .core.embedding import BaseEmbeddingModel
from .core.enumeration import Role
from .core.llm import BaseLLM
from .core.schema import Message
from .core.utils import singleton
from .core.vector_store import BaseVectorStore
from .mem_agent.retriever import ReMeRetriever
from .mem_agent.summarizer import ReMeSummarizer, PersonalSummarizer
@ -22,6 +23,7 @@ from .mem_tool import (
)
@singleton
class ReMe(Application):
"""Simplified ReMe application that auto-initializes the service context."""
@ -62,6 +64,10 @@ class ReMe(Application):
self.vector_store: BaseVectorStore = C.get_vector_store("default")
self.embedding_model: BaseEmbeddingModel = C.get_embedding_model("default")
@staticmethod
def get_llm(name: str) -> BaseLLM:
return C.get_llm(name)
@staticmethod
def _prepare_messages(messages: list[dict | Message], user_id: str, assistant_id: str):
if not messages: