refactor(llm): replace concurrency control with request rate limiting

This commit is contained in:
jinli.yl 2026-01-19 01:02:55 +08:00
parent 74c1386a69
commit e7a36067eb
2 changed files with 31 additions and 21 deletions

View file

@ -16,14 +16,14 @@ llm:
default:
backend: openai
model_name: qwen3-30b-a3b-instruct-2507
max_concurrency: 20
request_interval: 1
temperature: 0.0001
qwen3_max_instruct:
backend: openai
model_name: qwen3-max
# temperature: 0.6
max_concurrency: 20
request_interval: 1
embedding_model:
default:

View file

@ -17,24 +17,25 @@ 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 = 10, raise_exception: bool = False, max_concurrency: int | None = None, **kwargs):
def __init__(self, model_name: str, max_retries: int = 10, raise_exception: bool = False, request_interval: float = 0.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_concurrency: Maximum concurrent requests for async operations. If None, no concurrency limit is applied.
request_interval: Minimum time interval (in seconds) between consecutive requests. Default is 0.0 (no interval).
**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_concurrency: int | None = max_concurrency
self.request_interval: float = request_interval
self.kwargs: dict = kwargs
# Concurrency control for async operations
self._semaphore: asyncio.Semaphore | None = asyncio.Semaphore(max_concurrency) if max_concurrency else None
# Request rate control for async operations
self._last_request_time: float = 0.0
self._request_lock: asyncio.Lock = asyncio.Lock()
@staticmethod
def _accumulate_tool_call_chunk(tool_call, ret_tools: list[ToolCall]):
@ -128,14 +129,18 @@ class BaseLLM(ABC):
model_name: Optional model name to override self.model_name
**kwargs: Additional parameters
"""
# Apply concurrency control if configured
if self._semaphore:
async with self._semaphore:
async for chunk in self._stream_chat_impl(messages, tools, model_name, **kwargs):
yield chunk
else:
async for chunk in self._stream_chat_impl(messages, tools, model_name, **kwargs):
yield chunk
# Apply request rate limiting if configured
if self.request_interval > 0:
async with self._request_lock:
current_time = time.time()
elapsed = current_time - self._last_request_time
if elapsed < self.request_interval:
sleep_time = self.request_interval - elapsed
await asyncio.sleep(sleep_time)
self._last_request_time = time.time()
async for chunk in self._stream_chat_impl(messages, tools, model_name, **kwargs):
yield chunk
async def _stream_chat_impl(
self,
@ -356,12 +361,17 @@ class BaseLLM(ABC):
model_name: Optional model name to override self.model_name
**kwargs: Additional parameters
"""
# Apply concurrency control if configured
if self._semaphore:
async with self._semaphore:
return await self._chat_impl(messages, tools, enable_stream_print, callback_fn, default_value, model_name, **kwargs)
else:
return await self._chat_impl(messages, tools, enable_stream_print, callback_fn, default_value, model_name, **kwargs)
# Apply request rate limiting if configured
if self.request_interval > 0:
async with self._request_lock:
current_time = time.time()
elapsed = current_time - self._last_request_time
if elapsed < self.request_interval:
sleep_time = self.request_interval - elapsed
await asyncio.sleep(sleep_time)
self._last_request_time = time.time()
return await self._chat_impl(messages, tools, enable_stream_print, callback_fn, default_value, model_name, **kwargs)
async def _chat_impl(
self,