diff --git a/reme_ai/core/llm/base_llm.py b/reme_ai/core/llm/base_llm.py index e2c800e5..b5bfb4af 100644 --- a/reme_ai/core/llm/base_llm.py +++ b/reme_ai/core/llm/base_llm.py @@ -48,58 +48,60 @@ class BaseLLM(ABC): if self.max_rps is None: return - while True: - 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 have space in the rate limit window, proceed - if len(self._request_timestamps) < self.max_rps: - # Record this request - self._request_timestamps.append(current_time) - return - - # Calculate how long to wait - oldest_timestamp = self._request_timestamps[0] - wait_time = self.rps_window - (current_time - oldest_timestamp) + current_time = time.time() + + # Clean up old timestamps and calculate wait time in one critical section + async with self._rate_limit_lock: + # 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() - # Wait OUTSIDE the lock so other requests can proceed - 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) - # Loop back to re-acquire lock and check again + # If we have space in the rate limit window, record and proceed immediately + if len(self._request_timestamps) < self.max_rps: + self._request_timestamps.append(current_time) + return + + # Calculate how long to wait until the oldest request expires + oldest_timestamp = self._request_timestamps[0] + wait_time = oldest_timestamp + self.rps_window - current_time + + # Wait OUTSIDE the lock so other requests can proceed + 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 + 0.001) # Add small buffer to ensure timestamp expires + + # After waiting, recursively check again (oldest timestamp should now be expired) + await self._wait_for_rate_limit() 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 - while True: - 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 have space in the rate limit window, proceed - if len(self._request_timestamps) < self.max_rps: - # Record this request - self._request_timestamps.append(current_time) - return - - # Calculate how long to wait - oldest_timestamp = self._request_timestamps[0] - wait_time = self.rps_window - (current_time - oldest_timestamp) + current_time = time.time() + + # Clean up old timestamps and calculate wait time in one critical section + with self._rate_limit_lock_sync: + # 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() - # Wait OUTSIDE the lock so other requests can proceed - 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) - # Loop back to re-acquire lock and check again + # If we have space in the rate limit window, record and proceed immediately + if len(self._request_timestamps) < self.max_rps: + self._request_timestamps.append(current_time) + return + + # Calculate how long to wait until the oldest request expires + oldest_timestamp = self._request_timestamps[0] + wait_time = oldest_timestamp + self.rps_window - current_time + + # Wait OUTSIDE the lock so other requests can proceed + 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 + 0.001) # Add small buffer to ensure timestamp expires + + # After waiting, recursively check again (oldest timestamp should now be expired) + self._wait_for_rate_limit_sync() @staticmethod def _accumulate_tool_call_chunk(tool_call, ret_tools: list[ToolCall]):