mirror of https://github.com/aliasrobotics/cai.git
Merge 0ac1c715d0 into 62871b6f5a
This commit is contained in:
commit
1dbe7c2969
|
|
@ -77,6 +77,10 @@ from cai.errors import (
|
||||||
LLMTimeout,
|
LLMTimeout,
|
||||||
)
|
)
|
||||||
from cai.sdk.agents.models.chatcompletions.httpx_client import verbose_http_retries
|
from cai.sdk.agents.models.chatcompletions.httpx_client import verbose_http_retries
|
||||||
|
from cai.sdk.agents.models.chatcompletions.litellm_adapter import (
|
||||||
|
is_transient_litellm_provider_error,
|
||||||
|
provider_error_summary,
|
||||||
|
)
|
||||||
from cai.continuation import generate_continuation_advice, should_continue_automatically
|
from cai.continuation import generate_continuation_advice, should_continue_automatically
|
||||||
from litellm.exceptions import RateLimitError, Timeout
|
from litellm.exceptions import RateLimitError, Timeout
|
||||||
|
|
||||||
|
|
@ -1497,6 +1501,14 @@ def _run_streamed(agent, conversation_input, console, force_until_flag, ctf_glob
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
if isinstance(e, (LLMProviderUnavailable, LLMTimeout, LLMRateLimited)):
|
||||||
|
raise
|
||||||
|
if is_transient_litellm_provider_error(e):
|
||||||
|
summary = provider_error_summary(e)
|
||||||
|
logger.warning("Streaming provider error: %s", summary)
|
||||||
|
raise LLMProviderUnavailable(
|
||||||
|
f"Model provider disconnected during streaming: {summary}"
|
||||||
|
) from e
|
||||||
logger.error(f"Error occurred during streaming: {str(e)}", exc_info=True)
|
logger.error(f"Error occurred during streaming: {str(e)}", exc_info=True)
|
||||||
if _get_config().debug == 2:
|
if _get_config().debug == 2:
|
||||||
import traceback
|
import traceback
|
||||||
|
|
|
||||||
|
|
@ -11,6 +11,7 @@ from typing import List, Dict, Any, Optional
|
||||||
from rich.console import Console
|
from rich.console import Console
|
||||||
|
|
||||||
from cai.config import get_config
|
from cai.config import get_config
|
||||||
|
from cai.sdk.agents.models.chatcompletions.litellm_adapter import acompletion_with_timeout
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
@ -108,9 +109,7 @@ Generate a specific, actionable continuation prompt that:
|
||||||
IMPORTANT: Respond with ONLY the continuation prompt. No explanations, no "Here's a prompt:", just the direct instruction."""
|
IMPORTANT: Respond with ONLY the continuation prompt. No explanations, no "Here's a prompt:", just the direct instruction."""
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Use litellm directly, which is how the rest of the codebase handles API calls
|
# Use LiteLLM through CAI's timeout wrapper.
|
||||||
import litellm
|
|
||||||
|
|
||||||
# Enable debug logging for litellm if in debug mode
|
# Enable debug logging for litellm if in debug mode
|
||||||
if logger.isEnabledFor(logging.DEBUG):
|
if logger.isEnabledFor(logging.DEBUG):
|
||||||
logger.debug(f"Generating continuation advice with model: {model_name}")
|
logger.debug(f"Generating continuation advice with model: {model_name}")
|
||||||
|
|
@ -138,7 +137,7 @@ IMPORTANT: Respond with ONLY the continuation prompt. No explanations, no "Here'
|
||||||
|
|
||||||
# Make the API call
|
# Make the API call
|
||||||
logger.debug(f"Making API call with kwargs: {kwargs.get('model')}, provider: {kwargs.get('custom_llm_provider', 'default')}")
|
logger.debug(f"Making API call with kwargs: {kwargs.get('model')}, provider: {kwargs.get('custom_llm_provider', 'default')}")
|
||||||
response = await litellm.acompletion(**kwargs)
|
response = await acompletion_with_timeout(kwargs, stream=False, model_name=model_name)
|
||||||
|
|
||||||
# Extract content safely
|
# Extract content safely
|
||||||
continuation_prompt = None
|
continuation_prompt = None
|
||||||
|
|
|
||||||
|
|
@ -19,6 +19,8 @@ from pathlib import Path
|
||||||
from typing import Dict, List, Tuple, Optional
|
from typing import Dict, List, Tuple, Optional
|
||||||
from rich.console import Console
|
from rich.console import Console
|
||||||
|
|
||||||
|
from cai.sdk.agents.models.chatcompletions.litellm_adapter import acompletion_with_timeout
|
||||||
|
|
||||||
console = Console()
|
console = Console()
|
||||||
|
|
||||||
# Global cache for digest results (per-run caching)
|
# Global cache for digest results (per-run caching)
|
||||||
|
|
@ -519,7 +521,6 @@ OUTPUT REQUIREMENTS:
|
||||||
- Maximum 350 words"""
|
- Maximum 350 words"""
|
||||||
|
|
||||||
# Use LiteLLM for model compatibility (handles alias1, OpenRouter, etc.)
|
# Use LiteLLM for model compatibility (handles alias1, OpenRouter, etc.)
|
||||||
import litellm
|
|
||||||
|
|
||||||
model = os.getenv("CAI_CTR_DIGEST_MODEL", "alias1")
|
model = os.getenv("CAI_CTR_DIGEST_MODEL", "alias1")
|
||||||
|
|
||||||
|
|
@ -549,7 +550,7 @@ OUTPUT REQUIREMENTS:
|
||||||
kwargs["api_key"] = os.getenv("ALIAS_API_KEY", "sk-alias-1234567890")
|
kwargs["api_key"] = os.getenv("ALIAS_API_KEY", "sk-alias-1234567890")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
response = await litellm.acompletion(**kwargs)
|
response = await acompletion_with_timeout(kwargs, stream=False, model_name=model)
|
||||||
|
|
||||||
# Extract content (handle reasoning models like alias1/o1 that use reasoning_content)
|
# Extract content (handle reasoning models like alias1/o1 that use reasoning_content)
|
||||||
message = response.choices[0].message
|
message = response.choices[0].message
|
||||||
|
|
|
||||||
|
|
@ -8,6 +8,8 @@ Extracted from openai_chatcompletions.py [F] to reduce monolith size.
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import os
|
||||||
import time
|
import time
|
||||||
from typing import TYPE_CHECKING, Any, Literal, cast
|
from typing import TYPE_CHECKING, Any, Literal, cast
|
||||||
|
|
||||||
|
|
@ -15,6 +17,7 @@ import litellm
|
||||||
from openai import NOT_GIVEN, NotGiven
|
from openai import NOT_GIVEN, NotGiven
|
||||||
from openai.types.responses import Response
|
from openai.types.responses import Response
|
||||||
|
|
||||||
|
from cai.errors import LLMTimeout
|
||||||
from cai.util import get_ollama_api_base
|
from cai.util import get_ollama_api_base
|
||||||
from ..fake_id import FAKE_RESPONSES_ID
|
from ..fake_id import FAKE_RESPONSES_ID
|
||||||
|
|
||||||
|
|
@ -28,6 +31,140 @@ if TYPE_CHECKING:
|
||||||
from ...model_settings import ModelSettings
|
from ...model_settings import ModelSettings
|
||||||
|
|
||||||
|
|
||||||
|
_DEFAULT_MODEL_TIMEOUT = 180.0
|
||||||
|
|
||||||
|
|
||||||
|
def configured_model_timeout() -> float | None:
|
||||||
|
"""Return CAI's LiteLLM request timeout in seconds.
|
||||||
|
|
||||||
|
``CAI_MODEL_TIMEOUT`` is the public name. ``CAI_LLM_TIMEOUT`` is accepted
|
||||||
|
as a compatibility alias for local configs/scripts. Values <= 0 disable the
|
||||||
|
injected timeout and defer entirely to LiteLLM/provider defaults.
|
||||||
|
"""
|
||||||
|
raw = os.getenv("CAI_MODEL_TIMEOUT")
|
||||||
|
if raw is None:
|
||||||
|
raw = os.getenv("CAI_LLM_TIMEOUT")
|
||||||
|
if raw is None or raw == "":
|
||||||
|
return _DEFAULT_MODEL_TIMEOUT
|
||||||
|
try:
|
||||||
|
timeout = float(raw)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return _DEFAULT_MODEL_TIMEOUT
|
||||||
|
if timeout <= 0:
|
||||||
|
return None
|
||||||
|
return timeout
|
||||||
|
|
||||||
|
|
||||||
|
def apply_litellm_timeouts(kwargs: dict, *, stream: bool = False) -> dict:
|
||||||
|
"""Add bounded LiteLLM request timeouts unless the caller already set them."""
|
||||||
|
timeout = configured_model_timeout()
|
||||||
|
if timeout is None:
|
||||||
|
return kwargs
|
||||||
|
kwargs.setdefault("timeout", timeout)
|
||||||
|
if stream:
|
||||||
|
kwargs.setdefault("stream_timeout", timeout)
|
||||||
|
return kwargs
|
||||||
|
|
||||||
|
|
||||||
|
def is_transient_litellm_provider_error(exc: BaseException) -> bool:
|
||||||
|
"""Return True for provider/proxy failures that are safe to retry."""
|
||||||
|
return isinstance(
|
||||||
|
exc,
|
||||||
|
(
|
||||||
|
litellm.exceptions.APIConnectionError,
|
||||||
|
litellm.exceptions.BadGatewayError,
|
||||||
|
litellm.exceptions.InternalServerError,
|
||||||
|
litellm.exceptions.ServiceUnavailableError,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def provider_error_summary(exc: BaseException, *, limit: int = 220) -> str:
|
||||||
|
"""Compact provider error text for user-facing typed exceptions."""
|
||||||
|
message = " ".join(str(exc).split())
|
||||||
|
if len(message) > limit:
|
||||||
|
message = f"{message[:limit]}..."
|
||||||
|
return f"{type(exc).__name__}: {message}"
|
||||||
|
|
||||||
|
|
||||||
|
def _timeout_from_kwargs(kwargs: dict) -> float | None:
|
||||||
|
"""Return the effective numeric timeout for CAI's outer asyncio guard."""
|
||||||
|
raw_timeout = kwargs.get("timeout", configured_model_timeout())
|
||||||
|
if raw_timeout is None:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
timeout = float(raw_timeout)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return configured_model_timeout()
|
||||||
|
if timeout <= 0:
|
||||||
|
return None
|
||||||
|
return timeout
|
||||||
|
|
||||||
|
|
||||||
|
def wrap_stream_with_idle_timeout(stream_obj: Any, *, model_name: str, timeout: float | None = None) -> Any:
|
||||||
|
"""Bound waits for each streamed chunk.
|
||||||
|
|
||||||
|
Some LiteLLM/provider combinations return the stream object quickly, then
|
||||||
|
stall while the caller awaits the next SSE chunk. ``timeout``/
|
||||||
|
``stream_timeout`` do not consistently protect that phase, so CAI wraps the
|
||||||
|
async iterator itself. Non-async-iterable test doubles are returned as-is.
|
||||||
|
"""
|
||||||
|
if timeout is None:
|
||||||
|
timeout = configured_model_timeout()
|
||||||
|
if timeout is None or not hasattr(stream_obj, "__aiter__"):
|
||||||
|
return stream_obj
|
||||||
|
|
||||||
|
async def _iter_with_timeout():
|
||||||
|
iterator = stream_obj.__aiter__()
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
chunk = await asyncio.wait_for(iterator.__anext__(), timeout=timeout)
|
||||||
|
except StopAsyncIteration:
|
||||||
|
return
|
||||||
|
except asyncio.TimeoutError as exc:
|
||||||
|
raise LLMTimeout(
|
||||||
|
f"Timed out waiting for streamed model chunk after {timeout:g}s "
|
||||||
|
f"[{model_name}]"
|
||||||
|
) from exc
|
||||||
|
yield chunk
|
||||||
|
|
||||||
|
return _iter_with_timeout()
|
||||||
|
|
||||||
|
|
||||||
|
async def acompletion_with_timeout(
|
||||||
|
kwargs: dict,
|
||||||
|
*,
|
||||||
|
stream: bool = False,
|
||||||
|
model_name: str | None = None,
|
||||||
|
) -> Any:
|
||||||
|
"""Call LiteLLM with CAI request and stream-idle timeouts applied."""
|
||||||
|
kwargs = apply_litellm_timeouts(kwargs, stream=stream)
|
||||||
|
timeout = _timeout_from_kwargs(kwargs)
|
||||||
|
model_label = str(model_name or kwargs.get("model") or "unknown model")
|
||||||
|
|
||||||
|
completion_coro = litellm.acompletion(**kwargs)
|
||||||
|
try:
|
||||||
|
if timeout is None:
|
||||||
|
result = await completion_coro
|
||||||
|
else:
|
||||||
|
result = await asyncio.wait_for(completion_coro, timeout=timeout)
|
||||||
|
except asyncio.TimeoutError as exc:
|
||||||
|
raise LLMTimeout(
|
||||||
|
f"Timed out waiting for model response after {timeout:g}s [{model_label}]"
|
||||||
|
) from exc
|
||||||
|
|
||||||
|
if stream:
|
||||||
|
return wrap_stream_with_idle_timeout(result, model_name=model_label, timeout=timeout)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
# Backward-compatible private aliases for local/internal imports.
|
||||||
|
_configured_model_timeout = configured_model_timeout
|
||||||
|
_apply_litellm_timeouts = apply_litellm_timeouts
|
||||||
|
_wrap_stream_with_idle_timeout = wrap_stream_with_idle_timeout
|
||||||
|
_acompletion_with_timeout = acompletion_with_timeout
|
||||||
|
|
||||||
|
|
||||||
def _build_response_obj(
|
def _build_response_obj(
|
||||||
model: str,
|
model: str,
|
||||||
model_settings: "ModelSettings",
|
model_settings: "ModelSettings",
|
||||||
|
|
@ -66,13 +203,14 @@ async def fetch_response_litellm_openai(
|
||||||
too long, truncate all tool_call ids in the messages to 40 characters
|
too long, truncate all tool_call ids in the messages to 40 characters
|
||||||
and retry once silently.
|
and retry once silently.
|
||||||
"""
|
"""
|
||||||
|
kwargs = _apply_litellm_timeouts(kwargs, stream=stream)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if stream:
|
if stream:
|
||||||
ret = await litellm.acompletion(**kwargs)
|
stream_obj = await acompletion_with_timeout(kwargs, stream=True, model_name=model_name)
|
||||||
stream_obj = await litellm.acompletion(**kwargs)
|
|
||||||
return _build_response_obj(model_name, model_settings, tool_choice, parallel_tool_calls), stream_obj
|
return _build_response_obj(model_name, model_settings, tool_choice, parallel_tool_calls), stream_obj
|
||||||
else:
|
else:
|
||||||
return await litellm.acompletion(**kwargs)
|
return await acompletion_with_timeout(kwargs, stream=False, model_name=model_name)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
error_msg = str(e)
|
error_msg = str(e)
|
||||||
if (
|
if (
|
||||||
|
|
@ -102,11 +240,10 @@ async def fetch_response_litellm_openai(
|
||||||
kwargs["messages"] = messages
|
kwargs["messages"] = messages
|
||||||
|
|
||||||
if stream:
|
if stream:
|
||||||
ret = await litellm.acompletion(**kwargs)
|
stream_obj = await acompletion_with_timeout(kwargs, stream=True, model_name=model_name)
|
||||||
stream_obj = await litellm.acompletion(**kwargs)
|
|
||||||
return _build_response_obj(model_name, model_settings, tool_choice, parallel_tool_calls), stream_obj
|
return _build_response_obj(model_name, model_settings, tool_choice, parallel_tool_calls), stream_obj
|
||||||
else:
|
else:
|
||||||
return await litellm.acompletion(**kwargs)
|
return await acompletion_with_timeout(kwargs, stream=False, model_name=model_name)
|
||||||
else:
|
else:
|
||||||
raise
|
raise
|
||||||
|
|
||||||
|
|
@ -150,15 +287,17 @@ async def fetch_response_litellm_ollama(
|
||||||
|
|
||||||
api_base = get_ollama_api_base()
|
api_base = get_ollama_api_base()
|
||||||
|
|
||||||
|
ollama_kwargs = _apply_litellm_timeouts(ollama_kwargs, stream=stream)
|
||||||
|
|
||||||
|
call_kwargs = {
|
||||||
|
**ollama_kwargs,
|
||||||
|
"api_base": api_base,
|
||||||
|
"custom_llm_provider": "openai",
|
||||||
|
}
|
||||||
|
|
||||||
if stream:
|
if stream:
|
||||||
response = _build_response_obj(model_name, model_settings, tool_choice, parallel_tool_calls)
|
response = _build_response_obj(model_name, model_settings, tool_choice, parallel_tool_calls)
|
||||||
stream_obj = await litellm.acompletion(
|
stream_obj = await acompletion_with_timeout(call_kwargs, stream=True, model_name=model_name)
|
||||||
**ollama_kwargs, api_base=api_base, custom_llm_provider="openai"
|
|
||||||
)
|
|
||||||
return response, stream_obj
|
return response, stream_obj
|
||||||
else:
|
else:
|
||||||
return await litellm.acompletion(
|
return await acompletion_with_timeout(call_kwargs, stream=False, model_name=model_name)
|
||||||
**ollama_kwargs,
|
|
||||||
api_base=api_base,
|
|
||||||
custom_llm_provider="openai",
|
|
||||||
)
|
|
||||||
|
|
|
||||||
|
|
@ -130,8 +130,11 @@ from .chatcompletions.httpx_client import (
|
||||||
verbose_http_retries,
|
verbose_http_retries,
|
||||||
)
|
)
|
||||||
from .chatcompletions.litellm_adapter import (
|
from .chatcompletions.litellm_adapter import (
|
||||||
|
acompletion_with_timeout,
|
||||||
fetch_response_litellm_openai as _fetch_litellm_openai_impl,
|
fetch_response_litellm_openai as _fetch_litellm_openai_impl,
|
||||||
fetch_response_litellm_ollama as _fetch_litellm_ollama_impl,
|
fetch_response_litellm_ollama as _fetch_litellm_ollama_impl,
|
||||||
|
is_transient_litellm_provider_error,
|
||||||
|
provider_error_summary,
|
||||||
)
|
)
|
||||||
from .chatcompletions.model import (
|
from .chatcompletions.model import (
|
||||||
ACTIVE_MODEL_INSTANCES,
|
ACTIVE_MODEL_INSTANCES,
|
||||||
|
|
@ -153,7 +156,7 @@ from cai.util.llm_api_base import (
|
||||||
resolve_llm_openai_compatible_base,
|
resolve_llm_openai_compatible_base,
|
||||||
resolve_llm_openai_compatible_api_key,
|
resolve_llm_openai_compatible_api_key,
|
||||||
)
|
)
|
||||||
from cai.errors import LLMEmptyAssistantError, LLMRateLimited, LLMTimeout
|
from cai.errors import LLMEmptyAssistantError, LLMProviderUnavailable, LLMRateLimited, LLMTimeout
|
||||||
from cai.util.gateway_rate_limiter import (
|
from cai.util.gateway_rate_limiter import (
|
||||||
COMPLETION_BUDGET_TOKENS,
|
COMPLETION_BUDGET_TOKENS,
|
||||||
get_gateway_rate_limiter,
|
get_gateway_rate_limiter,
|
||||||
|
|
@ -959,9 +962,11 @@ class OpenAIChatCompletionsModel(Model):
|
||||||
return result
|
return result
|
||||||
|
|
||||||
except (
|
except (
|
||||||
|
litellm.exceptions.APIConnectionError,
|
||||||
litellm.exceptions.BadGatewayError,
|
litellm.exceptions.BadGatewayError,
|
||||||
litellm.exceptions.ServiceUnavailableError,
|
litellm.exceptions.ServiceUnavailableError,
|
||||||
litellm.exceptions.InternalServerError,
|
litellm.exceptions.InternalServerError,
|
||||||
|
LLMProviderUnavailable,
|
||||||
) as e:
|
) as e:
|
||||||
# Transient server errors (502, 503, 500): retry with backoff
|
# Transient server errors (502, 503, 500): retry with backoff
|
||||||
self.logger.warning(f"Server error (high-level recovery): {str(e)[:200]}")
|
self.logger.warning(f"Server error (high-level recovery): {str(e)[:200]}")
|
||||||
|
|
@ -976,8 +981,11 @@ class OpenAIChatCompletionsModel(Model):
|
||||||
if self._high_level_retry_count > 3:
|
if self._high_level_retry_count > 3:
|
||||||
self._high_level_retry_count = 0
|
self._high_level_retry_count = 0
|
||||||
if verbose_http_retries():
|
if verbose_http_retries():
|
||||||
print(f"\n❌ Server error after 3 recovery attempts [{self.model}]")
|
print(f"\n❌ Provider error after 3 recovery attempts [{self.model}]")
|
||||||
raise
|
raise LLMProviderUnavailable(
|
||||||
|
f"Provider unavailable after 3 recovery attempts "
|
||||||
|
f"[{self.model}]: {provider_error_summary(e)}"
|
||||||
|
) from e
|
||||||
|
|
||||||
wait_secs = 10 * self._high_level_retry_count # 10s, 20s, 30s
|
wait_secs = 10 * self._high_level_retry_count # 10s, 20s, 30s
|
||||||
self.logger.warning(
|
self.logger.warning(
|
||||||
|
|
@ -1784,8 +1792,13 @@ class OpenAIChatCompletionsModel(Model):
|
||||||
|
|
||||||
# Clean retry: same input, no "continue" in history
|
# Clean retry: same input, no "continue" in history
|
||||||
async for event in self.stream_response(
|
async for event in self.stream_response(
|
||||||
system_instructions, input, model_settings,
|
system_instructions,
|
||||||
tools, output_schema, handoffs, tracing,
|
input,
|
||||||
|
model_settings,
|
||||||
|
tools,
|
||||||
|
output_schema,
|
||||||
|
handoffs,
|
||||||
|
tracing,
|
||||||
):
|
):
|
||||||
yield event
|
yield event
|
||||||
self._high_level_retry_count = 0
|
self._high_level_retry_count = 0
|
||||||
|
|
@ -1813,8 +1826,58 @@ class OpenAIChatCompletionsModel(Model):
|
||||||
|
|
||||||
# Clean retry: same input, no "continue" in history
|
# Clean retry: same input, no "continue" in history
|
||||||
async for event in self.stream_response(
|
async for event in self.stream_response(
|
||||||
system_instructions, input, model_settings,
|
system_instructions,
|
||||||
tools, output_schema, handoffs, tracing,
|
input,
|
||||||
|
model_settings,
|
||||||
|
tools,
|
||||||
|
output_schema,
|
||||||
|
handoffs,
|
||||||
|
tracing,
|
||||||
|
):
|
||||||
|
yield event
|
||||||
|
self._high_level_retry_count = 0
|
||||||
|
return
|
||||||
|
|
||||||
|
except (
|
||||||
|
litellm.exceptions.APIConnectionError,
|
||||||
|
litellm.exceptions.BadGatewayError,
|
||||||
|
litellm.exceptions.ServiceUnavailableError,
|
||||||
|
litellm.exceptions.InternalServerError,
|
||||||
|
LLMProviderUnavailable,
|
||||||
|
) as e:
|
||||||
|
await stream_wait_hints.stop()
|
||||||
|
self.logger.warning(
|
||||||
|
"Transient provider error in stream_response [%s]: %s",
|
||||||
|
self.model,
|
||||||
|
provider_error_summary(e),
|
||||||
|
)
|
||||||
|
stop_active_timer()
|
||||||
|
start_idle_timer()
|
||||||
|
|
||||||
|
if not hasattr(self, "_high_level_retry_count"):
|
||||||
|
self._high_level_retry_count = 0
|
||||||
|
self._high_level_retry_count += 1
|
||||||
|
|
||||||
|
if self._high_level_retry_count > 3:
|
||||||
|
self._high_level_retry_count = 0
|
||||||
|
raise LLMProviderUnavailable(
|
||||||
|
f"Provider unavailable after 3 attempts "
|
||||||
|
f"[{self.model}]: {provider_error_summary(e)}"
|
||||||
|
) from e
|
||||||
|
|
||||||
|
await self._retry_with_backoff(
|
||||||
|
self._high_level_retry_count - 1, "Provider error"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Clean retry: same input, no "continue" in history
|
||||||
|
async for event in self.stream_response(
|
||||||
|
system_instructions,
|
||||||
|
input,
|
||||||
|
model_settings,
|
||||||
|
tools,
|
||||||
|
output_schema,
|
||||||
|
handoffs,
|
||||||
|
tracing,
|
||||||
):
|
):
|
||||||
yield event
|
yield event
|
||||||
self._high_level_retry_count = 0
|
self._high_level_retry_count = 0
|
||||||
|
|
@ -1994,10 +2057,7 @@ class OpenAIChatCompletionsModel(Model):
|
||||||
|
|
||||||
# Check if model supports reasoning (Claude or DeepSeek)
|
# Check if model supports reasoning (Claude or DeepSeek)
|
||||||
model_str_lower = str(self.model).lower()
|
model_str_lower = str(self.model).lower()
|
||||||
if (
|
if detect_claude_thinking_in_stream(str(self.model)):
|
||||||
detect_claude_thinking_in_stream(str(self.model))
|
|
||||||
or "deepseek" in model_str_lower
|
|
||||||
):
|
|
||||||
print_claude_reasoning_simple(
|
print_claude_reasoning_simple(
|
||||||
reasoning_content, self.agent_name, str(self.model)
|
reasoning_content, self.agent_name, str(self.model)
|
||||||
)
|
)
|
||||||
|
|
@ -3069,8 +3129,15 @@ class OpenAIChatCompletionsModel(Model):
|
||||||
raise
|
raise
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
# Handle other exceptions
|
# Provider/proxy errors are already retried above; keep logs concise.
|
||||||
logger.error(f"Error in stream_response: {e}")
|
if isinstance(e, (LLMProviderUnavailable, LLMTimeout, LLMRateLimited)) or (
|
||||||
|
is_transient_litellm_provider_error(e)
|
||||||
|
):
|
||||||
|
logger.warning(
|
||||||
|
"Model provider error in stream_response: %s", provider_error_summary(e)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.error(f"Error in stream_response: {e}")
|
||||||
raise
|
raise
|
||||||
|
|
||||||
finally:
|
finally:
|
||||||
|
|
@ -4119,10 +4186,14 @@ class OpenAIChatCompletionsModel(Model):
|
||||||
tools=[],
|
tools=[],
|
||||||
parallel_tool_calls=parallel_tool_calls or False,
|
parallel_tool_calls=parallel_tool_calls or False,
|
||||||
)
|
)
|
||||||
stream_obj = await litellm.acompletion(**retry_kwargs)
|
stream_obj = await acompletion_with_timeout(
|
||||||
|
retry_kwargs, stream=True, model_name=str(self.model)
|
||||||
|
)
|
||||||
return response, stream_obj
|
return response, stream_obj
|
||||||
else:
|
else:
|
||||||
ret = await litellm.acompletion(**retry_kwargs)
|
ret = await acompletion_with_timeout(
|
||||||
|
retry_kwargs, stream=False, model_name=str(self.model)
|
||||||
|
)
|
||||||
return ret
|
return ret
|
||||||
except Exception:
|
except Exception:
|
||||||
# If retry also fails, raise the original error
|
# If retry also fails, raise the original error
|
||||||
|
|
@ -4171,11 +4242,15 @@ class OpenAIChatCompletionsModel(Model):
|
||||||
tools=[],
|
tools=[],
|
||||||
parallel_tool_calls=parallel_tool_calls or False,
|
parallel_tool_calls=parallel_tool_calls or False,
|
||||||
)
|
)
|
||||||
stream_obj = await litellm.acompletion(**qwen_params)
|
stream_obj = await acompletion_with_timeout(
|
||||||
|
qwen_params, stream=True, model_name=str(self.model)
|
||||||
|
)
|
||||||
return response, stream_obj
|
return response, stream_obj
|
||||||
else:
|
else:
|
||||||
# Non-streaming case
|
# Non-streaming case
|
||||||
ret = await litellm.acompletion(**qwen_params)
|
ret = await acompletion_with_timeout(
|
||||||
|
qwen_params, stream=False, model_name=str(self.model)
|
||||||
|
)
|
||||||
return ret
|
return ret
|
||||||
except Exception as direct_e:
|
except Exception as direct_e:
|
||||||
# All approaches failed, log and raise the original error
|
# All approaches failed, log and raise the original error
|
||||||
|
|
|
||||||
|
|
@ -15,7 +15,7 @@ import json
|
||||||
import os
|
import os
|
||||||
from typing import Dict, List, Any, Optional
|
from typing import Dict, List, Any, Optional
|
||||||
from cai.agents.agent_builder import AgentBuilder
|
from cai.agents.agent_builder import AgentBuilder
|
||||||
import litellm
|
from cai.sdk.agents.models.chatcompletions.litellm_adapter import acompletion_with_timeout
|
||||||
|
|
||||||
|
|
||||||
class AgentCreationConfirmed(Message):
|
class AgentCreationConfirmed(Message):
|
||||||
|
|
@ -514,14 +514,18 @@ IMPORTANT: The "tools" field in your response must be exactly: {json.dumps(selec
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Use litellm to generate the configuration
|
# Use litellm to generate the configuration
|
||||||
response = await litellm.acompletion(
|
model_name = os.getenv("CAI_MODEL", "gpt-4")
|
||||||
model=os.getenv("CAI_MODEL", "gpt-4"),
|
kwargs = {
|
||||||
messages=[
|
"model": model_name,
|
||||||
|
"messages": [
|
||||||
{"role": "system", "content": "You are an AI agent configuration generator. Always respond with valid JSON only."},
|
{"role": "system", "content": "You are an AI agent configuration generator. Always respond with valid JSON only."},
|
||||||
{"role": "user", "content": meta_prompt}
|
{"role": "user", "content": meta_prompt}
|
||||||
],
|
],
|
||||||
temperature=0.7,
|
"temperature": 0.7,
|
||||||
max_tokens=2000
|
"max_tokens": 2000,
|
||||||
|
}
|
||||||
|
response = await acompletion_with_timeout(
|
||||||
|
kwargs, stream=False, model_name=model_name
|
||||||
)
|
)
|
||||||
|
|
||||||
# Parse the response
|
# Parse the response
|
||||||
|
|
|
||||||
|
|
@ -4208,7 +4208,7 @@ def create_claude_thinking_context(agent_name, counter, model):
|
||||||
context = {
|
context = {
|
||||||
"thinking_id": thinking_id,
|
"thinking_id": thinking_id,
|
||||||
"live": live,
|
"live": live,
|
||||||
"panel": panel,
|
"panel": None,
|
||||||
"header": header,
|
"header": header,
|
||||||
"thinking_content": thinking_content,
|
"thinking_content": thinking_content,
|
||||||
"timestamp": timestamp,
|
"timestamp": timestamp,
|
||||||
|
|
@ -4358,10 +4358,29 @@ def finish_claude_thinking_display(context):
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _raw_reasoning_display_enabled() -> bool:
|
||||||
|
"""Return whether raw provider reasoning text should be displayed.
|
||||||
|
|
||||||
|
DeepSeek-compatible gateways often stream ``reasoning_content`` in tiny
|
||||||
|
deltas. Rendering those deltas by default floods the terminal and can leak
|
||||||
|
model-internal scratch text. Keep it opt-in for DeepSeek via
|
||||||
|
``CAI_SHOW_REASONING=true`` (or ``CAI_SHOW_THINKING=true`` for older local
|
||||||
|
configs).
|
||||||
|
"""
|
||||||
|
raw = os.getenv("CAI_SHOW_REASONING")
|
||||||
|
if raw is None:
|
||||||
|
raw = os.getenv("CAI_SHOW_THINKING")
|
||||||
|
if raw is None:
|
||||||
|
return False
|
||||||
|
return raw.strip().lower() in ("1", "true", "yes", "on")
|
||||||
|
|
||||||
|
|
||||||
def detect_claude_thinking_in_stream(model_name):
|
def detect_claude_thinking_in_stream(model_name):
|
||||||
"""
|
"""
|
||||||
Detect if a model should show thinking/reasoning display.
|
Detect if a model should show thinking/reasoning display.
|
||||||
Applies to Claude and DeepSeek models with reasoning capability.
|
|
||||||
|
Claude keeps the historical default. DeepSeek raw reasoning is opt-in
|
||||||
|
because providers commonly stream it token-by-token as ``reasoning_content``.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
model_name: The model name to check
|
model_name: The model name to check
|
||||||
|
|
@ -4389,17 +4408,10 @@ def detect_claude_thinking_in_stream(model_name):
|
||||||
or "thinking" in model_str
|
or "thinking" in model_str
|
||||||
)
|
)
|
||||||
|
|
||||||
# Check for DeepSeek models with reasoning capability
|
# DeepSeek reasoning display is intentionally opt-in. The text is still
|
||||||
has_deepseek_reasoning = "deepseek" in model_str and (
|
# accumulated internally by the model stream handler so empty-response
|
||||||
# DeepSeek reasoner models
|
# detection and accounting keep working; it is just not printed by default.
|
||||||
"reasoner" in model_str
|
has_deepseek_reasoning = "deepseek" in model_str and _raw_reasoning_display_enabled()
|
||||||
or
|
|
||||||
# DeepSeek chat models also support reasoning
|
|
||||||
"chat" in model_str
|
|
||||||
or
|
|
||||||
# Generic deepseek models likely support it
|
|
||||||
"/" in model_str # e.g., deepseek/deepseek-chat
|
|
||||||
)
|
|
||||||
|
|
||||||
return has_claude_reasoning or has_deepseek_reasoning
|
return has_claude_reasoning or has_deepseek_reasoning
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -4,9 +4,11 @@ from types import SimpleNamespace
|
||||||
from unittest.mock import Mock
|
from unittest.mock import Mock
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
import litellm
|
||||||
|
|
||||||
from cai import cli_headless
|
from cai import cli_headless
|
||||||
from cai import parallel_worker
|
from cai import parallel_worker
|
||||||
|
from cai.errors import LLMProviderUnavailable
|
||||||
|
|
||||||
|
|
||||||
def test_non_streamed_cancelled_error_uses_interrupt_flow(monkeypatch):
|
def test_non_streamed_cancelled_error_uses_interrupt_flow(monkeypatch):
|
||||||
|
|
@ -41,6 +43,33 @@ def test_streamed_cancelled_error_uses_interrupt_flow(monkeypatch):
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_streamed_provider_disconnect_raises_typed_provider_error(monkeypatch):
|
||||||
|
class DummyResult:
|
||||||
|
async def stream_events(self):
|
||||||
|
raise litellm.exceptions.InternalServerError(
|
||||||
|
message="DeepseekException - Server disconnected",
|
||||||
|
llm_provider="deepseek",
|
||||||
|
model="deepseek/deepseek-v4-pro",
|
||||||
|
)
|
||||||
|
yield # pragma: no cover
|
||||||
|
|
||||||
|
def _cleanup_tasks(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
cli_headless.Runner, "run_streamed", lambda *_args, **_kwargs: DummyResult()
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(LLMProviderUnavailable, match="Server disconnected"):
|
||||||
|
cli_headless._run_streamed(
|
||||||
|
SimpleNamespace(model=SimpleNamespace(message_history=[])),
|
||||||
|
"input",
|
||||||
|
Mock(),
|
||||||
|
False,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_simple_parallel_cancelled_error_uses_interrupt_flow(monkeypatch):
|
def test_simple_parallel_cancelled_error_uses_interrupt_flow(monkeypatch):
|
||||||
dummy_agent = SimpleNamespace(model=SimpleNamespace(model="test-model", message_history=[]))
|
dummy_agent = SimpleNamespace(model=SimpleNamespace(model="test-model", message_history=[]))
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,7 @@
|
||||||
from collections.abc import AsyncIterator
|
from collections.abc import AsyncIterator
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
import litellm
|
||||||
from openai.types.chat.chat_completion_chunk import (
|
from openai.types.chat.chat_completion_chunk import (
|
||||||
ChatCompletionChunk,
|
ChatCompletionChunk,
|
||||||
Choice,
|
Choice,
|
||||||
|
|
@ -121,6 +122,66 @@ async def test_stream_response_yields_events_for_text_content(monkeypatch) -> No
|
||||||
assert completed_resp.usage.total_tokens == 12
|
assert completed_resp.usage.total_tokens == 12
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.allow_call_model_methods
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stream_response_retries_transient_provider_disconnect(monkeypatch) -> None:
|
||||||
|
chunk = ChatCompletionChunk(
|
||||||
|
id="chunk-id",
|
||||||
|
created=1,
|
||||||
|
model="fake",
|
||||||
|
object="chat.completion.chunk",
|
||||||
|
choices=[Choice(index=0, delta=ChoiceDelta(content="ok"))],
|
||||||
|
usage=CompletionUsage(completion_tokens=1, prompt_tokens=1, total_tokens=2),
|
||||||
|
)
|
||||||
|
|
||||||
|
async def fake_stream() -> AsyncIterator[ChatCompletionChunk]:
|
||||||
|
yield chunk
|
||||||
|
|
||||||
|
calls = {"count": 0}
|
||||||
|
|
||||||
|
async def patched_fetch_response(self, *args, **kwargs):
|
||||||
|
calls["count"] += 1
|
||||||
|
if calls["count"] == 1:
|
||||||
|
raise litellm.exceptions.InternalServerError(
|
||||||
|
message="DeepseekException - Server disconnected",
|
||||||
|
llm_provider="deepseek",
|
||||||
|
model="deepseek/deepseek-v4-pro",
|
||||||
|
)
|
||||||
|
resp = Response(
|
||||||
|
id="resp-id",
|
||||||
|
created_at=0,
|
||||||
|
model="fake-model",
|
||||||
|
object="response",
|
||||||
|
output=[],
|
||||||
|
tool_choice="none",
|
||||||
|
tools=[],
|
||||||
|
parallel_tool_calls=False,
|
||||||
|
)
|
||||||
|
return resp, fake_stream()
|
||||||
|
|
||||||
|
async def no_sleep_retry(self, *_args, **_kwargs):
|
||||||
|
return None
|
||||||
|
|
||||||
|
monkeypatch.setattr(OpenAIChatCompletionsModel, "_fetch_response", patched_fetch_response)
|
||||||
|
monkeypatch.setattr(OpenAIChatCompletionsModel, "_retry_with_backoff", no_sleep_retry)
|
||||||
|
|
||||||
|
model = OpenAIProvider(use_responses=False).get_model(cai_model)
|
||||||
|
events = []
|
||||||
|
async for event in model.stream_response(
|
||||||
|
system_instructions=None,
|
||||||
|
input="",
|
||||||
|
model_settings=ModelSettings(),
|
||||||
|
tools=[],
|
||||||
|
output_schema=None,
|
||||||
|
handoffs=[],
|
||||||
|
tracing=ModelTracing.DISABLED,
|
||||||
|
):
|
||||||
|
events.append(event)
|
||||||
|
|
||||||
|
assert calls["count"] == 2
|
||||||
|
assert any(getattr(event, "delta", None) == "ok" for event in events)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.allow_call_model_methods
|
@pytest.mark.allow_call_model_methods
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_stream_response_yields_events_for_refusal_content(monkeypatch) -> None:
|
async def test_stream_response_yields_events_for_refusal_content(monkeypatch) -> None:
|
||||||
|
|
|
||||||
|
|
@ -0,0 +1,213 @@
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from openai import NOT_GIVEN
|
||||||
|
|
||||||
|
from cai.errors import LLMTimeout
|
||||||
|
|
||||||
|
from cai.sdk.agents.model_settings import ModelSettings
|
||||||
|
from cai.sdk.agents.models.chatcompletions.litellm_adapter import (
|
||||||
|
fetch_response_litellm_openai,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_litellm_streaming_fetch_opens_one_completion(monkeypatch):
|
||||||
|
calls = []
|
||||||
|
sentinel_stream = object()
|
||||||
|
|
||||||
|
async def fake_acompletion(**kwargs):
|
||||||
|
calls.append(kwargs.copy())
|
||||||
|
return sentinel_stream
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"cai.sdk.agents.models.chatcompletions.litellm_adapter.litellm.acompletion",
|
||||||
|
fake_acompletion,
|
||||||
|
)
|
||||||
|
|
||||||
|
response, stream = await fetch_response_litellm_openai(
|
||||||
|
kwargs={"model": "deepseek/deepseek-v4-pro", "messages": [], "stream": True},
|
||||||
|
model_name="deepseek/deepseek-v4-pro",
|
||||||
|
model_settings=ModelSettings(),
|
||||||
|
tool_choice=NOT_GIVEN,
|
||||||
|
stream=True,
|
||||||
|
parallel_tool_calls=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert stream is sentinel_stream
|
||||||
|
assert response.model == "deepseek/deepseek-v4-pro"
|
||||||
|
assert len(calls) == 1
|
||||||
|
assert calls[0]["stream"] is True
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_litellm_streaming_tool_call_id_retry_opens_one_retry_stream(monkeypatch):
|
||||||
|
calls = []
|
||||||
|
sentinel_stream = object()
|
||||||
|
|
||||||
|
async def fake_acompletion(**kwargs):
|
||||||
|
calls.append(kwargs.copy())
|
||||||
|
if len(calls) == 1:
|
||||||
|
raise Exception("Invalid 'messages': tool_call_id string too long maximum length")
|
||||||
|
return sentinel_stream
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"cai.sdk.agents.models.chatcompletions.litellm_adapter.litellm.acompletion",
|
||||||
|
fake_acompletion,
|
||||||
|
)
|
||||||
|
|
||||||
|
long_id = "call_" + "x" * 80
|
||||||
|
kwargs = {
|
||||||
|
"model": "deepseek/deepseek-v4-pro",
|
||||||
|
"messages": [
|
||||||
|
{"role": "tool", "tool_call_id": long_id, "content": "ok"},
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"tool_calls": [
|
||||||
|
{
|
||||||
|
"id": long_id,
|
||||||
|
"type": "function",
|
||||||
|
"function": {"name": "probe", "arguments": "{}"},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
},
|
||||||
|
],
|
||||||
|
"stream": True,
|
||||||
|
}
|
||||||
|
|
||||||
|
_response, stream = await fetch_response_litellm_openai(
|
||||||
|
kwargs=kwargs,
|
||||||
|
model_name="deepseek/deepseek-v4-pro",
|
||||||
|
model_settings=ModelSettings(),
|
||||||
|
tool_choice=NOT_GIVEN,
|
||||||
|
stream=True,
|
||||||
|
parallel_tool_calls=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert stream is sentinel_stream
|
||||||
|
assert len(calls) == 2
|
||||||
|
assert kwargs["messages"][0]["tool_call_id"] == long_id[:40]
|
||||||
|
assert kwargs["messages"][1]["tool_calls"][0]["id"] == long_id[:40]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_litellm_streaming_applies_default_model_timeout(monkeypatch):
|
||||||
|
monkeypatch.delenv("CAI_MODEL_TIMEOUT", raising=False)
|
||||||
|
monkeypatch.delenv("CAI_LLM_TIMEOUT", raising=False)
|
||||||
|
calls = []
|
||||||
|
sentinel_stream = object()
|
||||||
|
|
||||||
|
async def fake_acompletion(**kwargs):
|
||||||
|
calls.append(kwargs.copy())
|
||||||
|
return sentinel_stream
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"cai.sdk.agents.models.chatcompletions.litellm_adapter.litellm.acompletion",
|
||||||
|
fake_acompletion,
|
||||||
|
)
|
||||||
|
|
||||||
|
_response, stream = await fetch_response_litellm_openai(
|
||||||
|
kwargs={"model": "deepseek/deepseek-v4-pro", "messages": [], "stream": True},
|
||||||
|
model_name="deepseek/deepseek-v4-pro",
|
||||||
|
model_settings=ModelSettings(),
|
||||||
|
tool_choice=NOT_GIVEN,
|
||||||
|
stream=True,
|
||||||
|
parallel_tool_calls=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert stream is sentinel_stream
|
||||||
|
assert calls[0]["timeout"] == 180.0
|
||||||
|
assert calls[0]["stream_timeout"] == 180.0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_litellm_model_timeout_uses_env_override(monkeypatch):
|
||||||
|
monkeypatch.setenv("CAI_MODEL_TIMEOUT", "45")
|
||||||
|
calls = []
|
||||||
|
sentinel_response = object()
|
||||||
|
|
||||||
|
async def fake_acompletion(**kwargs):
|
||||||
|
calls.append(kwargs.copy())
|
||||||
|
return sentinel_response
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"cai.sdk.agents.models.chatcompletions.litellm_adapter.litellm.acompletion",
|
||||||
|
fake_acompletion,
|
||||||
|
)
|
||||||
|
|
||||||
|
response = await fetch_response_litellm_openai(
|
||||||
|
kwargs={"model": "deepseek/deepseek-v4-pro", "messages": [], "stream": False},
|
||||||
|
model_name="deepseek/deepseek-v4-pro",
|
||||||
|
model_settings=ModelSettings(),
|
||||||
|
tool_choice=NOT_GIVEN,
|
||||||
|
stream=False,
|
||||||
|
parallel_tool_calls=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response is sentinel_response
|
||||||
|
assert calls[0]["timeout"] == 45.0
|
||||||
|
assert "stream_timeout" not in calls[0]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_litellm_nonstream_times_out_when_completion_call_stalls(monkeypatch):
|
||||||
|
monkeypatch.setenv("CAI_MODEL_TIMEOUT", "0.01")
|
||||||
|
|
||||||
|
async def fake_acompletion(**_kwargs):
|
||||||
|
await asyncio.sleep(60)
|
||||||
|
return object()
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"cai.sdk.agents.models.chatcompletions.litellm_adapter.litellm.acompletion",
|
||||||
|
fake_acompletion,
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(LLMTimeout, match="Timed out waiting for model response"):
|
||||||
|
await asyncio.wait_for(
|
||||||
|
fetch_response_litellm_openai(
|
||||||
|
kwargs={
|
||||||
|
"model": "deepseek/deepseek-v4-pro",
|
||||||
|
"messages": [],
|
||||||
|
"stream": False,
|
||||||
|
},
|
||||||
|
model_name="deepseek/deepseek-v4-pro",
|
||||||
|
model_settings=ModelSettings(),
|
||||||
|
tool_choice=NOT_GIVEN,
|
||||||
|
stream=False,
|
||||||
|
parallel_tool_calls=False,
|
||||||
|
),
|
||||||
|
timeout=0.2,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_litellm_streaming_times_out_when_next_chunk_stalls(monkeypatch):
|
||||||
|
monkeypatch.setenv("CAI_MODEL_TIMEOUT", "0.01")
|
||||||
|
|
||||||
|
class StalledStream:
|
||||||
|
def __aiter__(self):
|
||||||
|
return self
|
||||||
|
|
||||||
|
async def __anext__(self):
|
||||||
|
await asyncio.sleep(60)
|
||||||
|
return object()
|
||||||
|
|
||||||
|
async def fake_acompletion(**_kwargs):
|
||||||
|
return StalledStream()
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"cai.sdk.agents.models.chatcompletions.litellm_adapter.litellm.acompletion",
|
||||||
|
fake_acompletion,
|
||||||
|
)
|
||||||
|
|
||||||
|
_response, stream = await fetch_response_litellm_openai(
|
||||||
|
kwargs={"model": "deepseek/deepseek-v4-pro", "messages": [], "stream": True},
|
||||||
|
model_name="deepseek/deepseek-v4-pro",
|
||||||
|
model_settings=ModelSettings(),
|
||||||
|
tool_choice=NOT_GIVEN,
|
||||||
|
stream=True,
|
||||||
|
parallel_tool_calls=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(LLMTimeout, match="Timed out waiting for streamed model chunk"):
|
||||||
|
await stream.__anext__()
|
||||||
|
|
@ -0,0 +1,39 @@
|
||||||
|
from cai.util import streaming
|
||||||
|
|
||||||
|
|
||||||
|
def test_create_claude_thinking_context_for_deepseek_does_not_reference_missing_panel(capsys):
|
||||||
|
streaming._CLAUDE_THINKING_PANELS.clear()
|
||||||
|
|
||||||
|
context = streaming.create_claude_thinking_context(
|
||||||
|
"Web App Pentester",
|
||||||
|
1,
|
||||||
|
"deepseek/deepseek-v4-pro",
|
||||||
|
)
|
||||||
|
|
||||||
|
captured = capsys.readouterr()
|
||||||
|
assert context is not None
|
||||||
|
assert "Error creating DeepSeek thinking context" not in captured.out
|
||||||
|
assert context["model_display"] == "DeepSeek"
|
||||||
|
assert context["accumulated_thinking"] == ""
|
||||||
|
assert context["is_started"] is False
|
||||||
|
|
||||||
|
thinking_id = context["thinking_id"]
|
||||||
|
assert streaming._CLAUDE_THINKING_PANELS[thinking_id] is context
|
||||||
|
streaming._CLAUDE_THINKING_PANELS.clear()
|
||||||
|
|
||||||
|
|
||||||
|
def test_deepseek_reasoning_display_is_opt_in(monkeypatch):
|
||||||
|
monkeypatch.delenv("CAI_SHOW_REASONING", raising=False)
|
||||||
|
monkeypatch.delenv("CAI_SHOW_THINKING", raising=False)
|
||||||
|
|
||||||
|
assert streaming.detect_claude_thinking_in_stream("deepseek/deepseek-v4-pro") is False
|
||||||
|
|
||||||
|
monkeypatch.setenv("CAI_SHOW_REASONING", "true")
|
||||||
|
assert streaming.detect_claude_thinking_in_stream("deepseek/deepseek-v4-pro") is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_claude_reasoning_display_keeps_historical_default(monkeypatch):
|
||||||
|
monkeypatch.delenv("CAI_SHOW_REASONING", raising=False)
|
||||||
|
monkeypatch.delenv("CAI_SHOW_THINKING", raising=False)
|
||||||
|
|
||||||
|
assert streaming.detect_claude_thinking_in_stream("claude-sonnet-4-20250514") is True
|
||||||
Loading…
Reference in New Issue