mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
feat(rust/messages): resend once without thinking after an invalid thinking signature
When Anthropic rejects replayed history with a 400 about an invalid or missing thinking signature, or an empty thinking block, the native route now strips the thinking and redacted_thinking blocks plus the top-level thinking param and sends the request once more, matching the python handler. Any other error, or a second failure, surfaces unchanged A new route-neutral RequestResent machine event lets the legacy logging adapter report the body that actually produced the response
This commit is contained in:
parent
0861f60c34
commit
ef42347675
9 changed files with 394 additions and 48 deletions
|
|
@ -367,6 +367,15 @@ impl PythonLifecycle for LegacyLogging {
|
|||
fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult<LifecycleStep> {
|
||||
match event {
|
||||
LifecycleEvent::Started { .. } => Ok(LifecycleStep::Done),
|
||||
LifecycleEvent::Machine(MachineEvent::RequestResent { body }) => {
|
||||
self.body = Some(
|
||||
to_py(py, body)?
|
||||
.into_bound(py)
|
||||
.cast_into::<PyDict>()?
|
||||
.unbind(),
|
||||
);
|
||||
Ok(LifecycleStep::Done)
|
||||
}
|
||||
LifecycleEvent::Machine(MachineEvent::ResponseReceived { raw }) => {
|
||||
let api_key = self
|
||||
.context
|
||||
|
|
@ -492,9 +501,7 @@ mod deployment_hooks_tests {
|
|||
|
||||
use litellm_host::event::{FailureOrigin, Timing};
|
||||
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle};
|
||||
use pyo3::exceptions::asyncio::CancelledError;
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::PyDict;
|
||||
use pyo3::{exceptions::asyncio::CancelledError, prelude::*, types::PyDict};
|
||||
use rstest::rstest;
|
||||
|
||||
use super::LegacyLogging;
|
||||
|
|
@ -787,8 +794,10 @@ mod payload_tests {
|
|||
use serde_json::{Map, Value, json};
|
||||
|
||||
use super::LegacyLogging;
|
||||
use crate::PythonLogger;
|
||||
use crate::test_support::{legacy_call, local, namespace, run};
|
||||
use crate::{
|
||||
PythonLogger,
|
||||
test_support::{legacy_call, local, namespace, run},
|
||||
};
|
||||
|
||||
/// The payload phases of `Logging` on top of `StubLogger`, with `pre_call` handing the
|
||||
/// payload to the case's `on_pre_call`.
|
||||
|
|
@ -837,7 +846,7 @@ check = lambda: None
|
|||
body: Value,
|
||||
secret_fields: &[&str],
|
||||
) -> WireRequest {
|
||||
before_send_bound(&[], script, optional_params, body, secret_fields)
|
||||
before_send_bound(&[], script, optional_params, body, secret_fields, None)
|
||||
}
|
||||
|
||||
/// [`before_send_with_secrets`] with `bindings` placed in the namespace before `script` runs.
|
||||
|
|
@ -847,6 +856,7 @@ check = lambda: None
|
|||
optional_params: Value,
|
||||
body: Value,
|
||||
secret_fields: &[&str],
|
||||
resent: Option<Value>,
|
||||
) -> WireRequest {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
|
|
@ -872,6 +882,13 @@ check = lambda: None
|
|||
body,
|
||||
};
|
||||
let step = logging.before_send(py, Box::new(wire), &context).unwrap();
|
||||
if let Some(body) = resent {
|
||||
let resent = MachineEvent::RequestResent { body };
|
||||
assert!(matches!(
|
||||
logging.emit(py, LifecycleEvent::Machine(&resent)).unwrap(),
|
||||
LifecycleStep::Done
|
||||
));
|
||||
}
|
||||
let raw = MachineEvent::ResponseReceived {
|
||||
raw: RawResponse {
|
||||
body: "raw response".into(),
|
||||
|
|
@ -1127,6 +1144,23 @@ def check():
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn post_call_logs_the_body_of_a_resent_request() {
|
||||
before_send_bound(
|
||||
&[],
|
||||
c"
|
||||
def check():
|
||||
_, _, additional_args = logger.post
|
||||
assert additional_args['complete_input_dict'] == {'messages': []}, additional_args
|
||||
assert logger.pre['complete_input_dict'] == {'messages': [{'role': 'user'}], 'thinking': {}}, logger.pre
|
||||
",
|
||||
json!({}),
|
||||
json!({"messages": [{"role": "user"}], "thinking": {}}),
|
||||
&[],
|
||||
Some(json!({"messages": []})),
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_request_runs_the_full_pre_call_and_post_call() {
|
||||
let wire = before_send(
|
||||
|
|
@ -1293,6 +1327,7 @@ def check():
|
|||
json!({}),
|
||||
Value::Object(body.clone()),
|
||||
&[],
|
||||
None,
|
||||
);
|
||||
|
||||
prop_assert_eq!(wire.body, edit.sent(&body));
|
||||
|
|
@ -1307,15 +1342,18 @@ mod terminal_tests {
|
|||
|
||||
use litellm_host::event::{FailureOrigin, Timing};
|
||||
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle};
|
||||
use pyo3::exceptions::PyRuntimeError;
|
||||
use pyo3::exceptions::asyncio::CancelledError;
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::PyDict;
|
||||
use pyo3::{
|
||||
exceptions::{PyRuntimeError, asyncio::CancelledError},
|
||||
prelude::*,
|
||||
types::PyDict,
|
||||
};
|
||||
use rstest::rstest;
|
||||
|
||||
use super::LegacyLogging;
|
||||
use crate::PythonLogger;
|
||||
use crate::test_support::{legacy_call, local, namespace, run};
|
||||
use crate::{
|
||||
PythonLogger,
|
||||
test_support::{legacy_call, local, namespace, run},
|
||||
};
|
||||
|
||||
const TIMING: Timing = Timing {
|
||||
start_time: 0.0,
|
||||
|
|
|
|||
|
|
@ -28,13 +28,17 @@ pub(super) async fn send(
|
|||
http_request(builder).await.map_err(network)
|
||||
}
|
||||
|
||||
pub(super) fn http_error(status: u16, body: &str) -> Error {
|
||||
Error::Transport(TransportError::Http {
|
||||
status,
|
||||
body: truncate_error_body(body),
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn provider_error(response: reqwest::Response) -> Error {
|
||||
let status = response.status().as_u16();
|
||||
match response.text().await {
|
||||
Ok(text) => Error::Transport(TransportError::Http {
|
||||
status,
|
||||
body: truncate_error_body(&text),
|
||||
}),
|
||||
Ok(text) => http_error(status, &text),
|
||||
Err(error) => network(error),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -12,9 +12,12 @@ use litellm_host::{
|
|||
machine::{HostChannel, MachineFault, RouteMachine},
|
||||
route::Route,
|
||||
};
|
||||
use litellm_llms::base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use litellm_types::{
|
||||
llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse,
|
||||
llms::anthropic_messages::{
|
||||
anthropic_request::AnthropicMessagesRequest, anthropic_response::AnthropicMessagesResponse,
|
||||
},
|
||||
utils::ProviderSpecificHeaders,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
|
|
@ -22,7 +25,7 @@ use serde_json::{Map, Value};
|
|||
use super::{
|
||||
Error,
|
||||
common_utils::messages_provider_config,
|
||||
handler::{decode_response, network, provider_error, send},
|
||||
handler::{decode_response, http_error, network, provider_error, send},
|
||||
prepare::{prepare_provider_request, resolve_provider},
|
||||
types::{MessagesRequest, MessagesShaping},
|
||||
};
|
||||
|
|
@ -183,9 +186,11 @@ async fn execute(
|
|||
)
|
||||
.await?;
|
||||
let response = send(&wire.url, &wire.headers, &wire.body, request.timeout).await?;
|
||||
if !response.status().is_success() {
|
||||
return Err(provider_error(response).await);
|
||||
}
|
||||
let response = if response.status().is_success() {
|
||||
response
|
||||
} else {
|
||||
resend_after_error(&host, request.config, &wire, request.timeout, response).await?
|
||||
};
|
||||
if stream {
|
||||
return relay(&host, response).await;
|
||||
}
|
||||
|
|
@ -198,6 +203,32 @@ async fn execute(
|
|||
.map(|message| MessagesOutput::Message(Box::new(message)))
|
||||
}
|
||||
|
||||
async fn resend_after_error(
|
||||
host: &MessagesHost,
|
||||
config: &dyn BaseAnthropicMessagesConfig,
|
||||
wire: &WireRequest,
|
||||
timeout: Option<Duration>,
|
||||
response: reqwest::Response,
|
||||
) -> Result<reqwest::Response, Error> {
|
||||
let status = response.status().as_u16();
|
||||
let text = response.text().await.map_err(network)?;
|
||||
let Some(request) = serde_json::from_value::<AnthropicMessagesRequest>(wire.body.clone())
|
||||
.ok()
|
||||
.and_then(|request| config.request_after_http_error(status, &text, request))
|
||||
else {
|
||||
return Err(http_error(status, &text));
|
||||
};
|
||||
let body = serde_json::to_value(request)
|
||||
.map_err(|error| Error::InvalidRequest(format!("invalid messages request: {error}")))?;
|
||||
host.emit(MachineEvent::RequestResent { body: body.clone() })
|
||||
.await?;
|
||||
let response = send(&wire.url, &wire.headers, &body, timeout).await?;
|
||||
if response.status().is_success() {
|
||||
return Ok(response);
|
||||
}
|
||||
Err(provider_error(response).await)
|
||||
}
|
||||
|
||||
/// Hands each upstream chunk to the caller as it arrives. A caller that stops reading
|
||||
/// ends the upstream read, and the call completes with what it delivered.
|
||||
async fn relay(
|
||||
|
|
|
|||
|
|
@ -410,9 +410,9 @@ mod azure_ai_tests {
|
|||
use litellm_llms::base_llm::ocr::error::Error;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use crate::ocr::route::LocalOcrHost;
|
||||
use crate::ocr::test_support::{
|
||||
MockResponse, mock_server, perform_ocr, perform_ocr_with, wire_request,
|
||||
use crate::ocr::{
|
||||
route::LocalOcrHost,
|
||||
test_support::{MockResponse, mock_server, perform_ocr, perform_ocr_with, wire_request},
|
||||
};
|
||||
|
||||
#[tokio::test]
|
||||
|
|
@ -712,8 +712,8 @@ mod azure_document_intelligence_tests {
|
|||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use crate::ocr::route::LocalOcrHost;
|
||||
use crate::ocr::{
|
||||
route::LocalOcrHost,
|
||||
test_support::{
|
||||
MockResponse, mock_server, ocr_client, perform_ocr, perform_ocr_with, wire_request,
|
||||
},
|
||||
|
|
@ -1715,9 +1715,9 @@ mod reducto_tests {
|
|||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use crate::ocr::route::LocalOcrHost;
|
||||
use crate::ocr::test_support::{
|
||||
MockResponse, mock_server, perform_ocr, perform_ocr_with, wire_request,
|
||||
use crate::ocr::{
|
||||
route::LocalOcrHost,
|
||||
test_support::{MockResponse, mock_server, perform_ocr, perform_ocr_with, wire_request},
|
||||
};
|
||||
|
||||
fn request_body(request: &str) -> Value {
|
||||
|
|
@ -2776,8 +2776,8 @@ pub(crate) mod tests {
|
|||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use crate::ocr::route::{LocalOcrHost, OcrOp, OcrOpResult, ocr_machine};
|
||||
use crate::ocr::{
|
||||
route::{LocalOcrHost, OcrOp, OcrOpResult, ocr_machine},
|
||||
test_support::{
|
||||
MockResponse, mock_server, ocr_client, perform_ocr, perform_ocr_with, wire_request,
|
||||
},
|
||||
|
|
@ -3059,6 +3059,7 @@ pub(crate) mod tests {
|
|||
match event {
|
||||
CallEvent::Started { .. } => "started",
|
||||
CallEvent::Machine(MachineEvent::ResponseReceived { .. }) => "response",
|
||||
CallEvent::Machine(MachineEvent::RequestResent { .. }) => "resent",
|
||||
CallEvent::Succeeded { .. } => "success",
|
||||
CallEvent::Failed { .. } => "failure",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,22 +1,27 @@
|
|||
use std::sync::Arc;
|
||||
use std::task::Poll;
|
||||
use std::{sync::Arc, task::Poll};
|
||||
|
||||
use futures_util::future::{AbortHandle, Abortable};
|
||||
use litellm_host::event::{FailureOrigin, Timing, epoch_seconds};
|
||||
use litellm_host::host::{Demand, HostOp, HostResult, HostStep};
|
||||
use litellm_host::machine::{HostFailure, Machine, MachineStep};
|
||||
use litellm_host::route::Route;
|
||||
use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError};
|
||||
use pyo3::gc::{PyTraverseError, PyVisit};
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::PyDict;
|
||||
use litellm_host::{
|
||||
event::{FailureOrigin, Timing, epoch_seconds},
|
||||
host::{Demand, HostOp, HostResult, HostStep},
|
||||
machine::{HostFailure, Machine, MachineStep},
|
||||
route::Route,
|
||||
};
|
||||
use pyo3::{
|
||||
exceptions::{PyBaseException, PyException, PyRuntimeError},
|
||||
gc::{PyTraverseError, PyVisit},
|
||||
prelude::*,
|
||||
types::PyDict,
|
||||
};
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
use crate::adapter::{
|
||||
InvokeError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state,
|
||||
use crate::{
|
||||
adapter::{
|
||||
InvokeError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state,
|
||||
},
|
||||
execution::{poll_async_value, run_async_value, run_sync_value},
|
||||
handle::{Execution, ExecutionBody, ExecutionStep},
|
||||
};
|
||||
use crate::execution::{poll_async_value, run_async_value, run_sync_value};
|
||||
use crate::handle::{Execution, ExecutionBody, ExecutionStep};
|
||||
|
||||
type RouteOf<H> = <H as RouteHost>::Route;
|
||||
type ErrorOf<H> = <RouteOf<H> as Route>::Error;
|
||||
|
|
@ -528,10 +533,14 @@ where
|
|||
mod tests {
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use litellm_host::event::{MachineEvent, RequestContext, WireRequest};
|
||||
use litellm_host::machine::{Interrupted, Step};
|
||||
use pyo3::exceptions::{PyBaseException, PyValueError};
|
||||
use pyo3::types::PyDict;
|
||||
use litellm_host::{
|
||||
event::{MachineEvent, RequestContext, WireRequest},
|
||||
machine::{Interrupted, Step},
|
||||
};
|
||||
use pyo3::{
|
||||
exceptions::{PyBaseException, PyValueError},
|
||||
types::PyDict,
|
||||
};
|
||||
|
||||
use super::*;
|
||||
|
||||
|
|
@ -793,6 +802,9 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
|
|||
LifecycleEvent::Machine(MachineEvent::ResponseReceived { raw }) => {
|
||||
format!("response:{}", raw.body)
|
||||
}
|
||||
LifecycleEvent::Machine(MachineEvent::RequestResent { body }) => {
|
||||
format!("resent:{body}")
|
||||
}
|
||||
LifecycleEvent::Succeeded { response, .. } => {
|
||||
format!("succeeded:{}", response.bind(py))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -55,6 +55,7 @@ pub enum FailureOrigin {
|
|||
/// What a machine reports while it runs.
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum MachineEvent {
|
||||
RequestResent { body: Value },
|
||||
ResponseReceived { raw: RawResponse },
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -232,6 +232,20 @@ fn retain_blocks(
|
|||
.collect()
|
||||
}
|
||||
|
||||
pub fn is_anthropic_invalid_thinking_block_error(error_text: &str) -> bool {
|
||||
let lower = error_text.to_lowercase();
|
||||
lower.contains("thinking")
|
||||
&& ((lower.contains("signature")
|
||||
&& (lower.contains("invalid") || lower.contains("valid string")))
|
||||
|| lower.contains("must contain thinking"))
|
||||
}
|
||||
|
||||
pub fn strip_thinking_blocks(messages: Vec<AnthropicMessage>) -> Vec<AnthropicMessage> {
|
||||
retain_blocks(messages, |block| {
|
||||
!block.is_type("thinking") && !block.is_type("redacted_thinking")
|
||||
})
|
||||
}
|
||||
|
||||
pub fn strip_empty_content_blocks(messages: Vec<AnthropicMessage>) -> Vec<AnthropicMessage> {
|
||||
retain_blocks(messages, |block| {
|
||||
!is_empty_text_block(block) && !is_empty_thinking_block(block)
|
||||
|
|
@ -1221,6 +1235,73 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic_invalid_signature(
|
||||
r#"{"type":"error","error":{"type":"invalid_request_error","message":"messages.3.content.3: Invalid `signature` in `thinking` block"},"request_id":"req_1"}"#
|
||||
)]
|
||||
#[case::bedrock_signature_not_a_string(
|
||||
r#"{"message":"messages.2.content.0.thinking.signature.str: Input should be a valid string"}"#
|
||||
)]
|
||||
#[case::vertex_signature_not_a_string(
|
||||
"messages.4.content.1.thinking.signature.str: Input should be a valid string"
|
||||
)]
|
||||
#[case::empty_thinking_text(
|
||||
r#"{"type":"error","error":{"type":"invalid_request_error","message":"messages.1.content.0.thinking: each thinking block must contain thinking"}}"#
|
||||
)]
|
||||
#[case::shouting("MESSAGES.0.CONTENT.0: INVALID `SIGNATURE` IN `THINKING` BLOCK")]
|
||||
fn invalid_thinking_block_errors_are_recognized(#[case] error_text: &str) {
|
||||
assert!(is_anthropic_invalid_thinking_block_error(error_text));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::empty("")]
|
||||
#[case::rate_limit("rate limit exceeded")]
|
||||
#[case::unrelated_invalid_request("invalid_request_error: model not found")]
|
||||
#[case::signature_without_invalid("thinking signature is malformed")]
|
||||
#[case::invalid_signature_without_thinking(
|
||||
"messages.0.content.0: Invalid `signature` in `tool_use` block"
|
||||
)]
|
||||
#[case::must_contain_without_thinking("each text block must contain text")]
|
||||
fn other_errors_are_not_invalid_thinking_block_errors(#[case] error_text: &str) {
|
||||
assert!(!is_anthropic_invalid_thinking_block_error(error_text));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::thinking_beside_text(
|
||||
json!([
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "thinking", "thinking": "plan", "signature": "sig"},
|
||||
{"type": "text", "text": "hello"}
|
||||
]}
|
||||
]),
|
||||
json!([
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": "hello"}]}
|
||||
])
|
||||
)]
|
||||
#[case::redacted_thinking_beside_tool_use(
|
||||
json!([{"role": "assistant", "content": [
|
||||
{"type": "redacted_thinking", "data": "opaque"},
|
||||
{"type": "tool_use", "id": "t1", "name": "Bash", "input": {}}
|
||||
]}]),
|
||||
json!([{"role": "assistant", "content": [{"type": "tool_use", "id": "t1", "name": "Bash", "input": {}}]}])
|
||||
)]
|
||||
#[case::message_left_without_blocks_is_dropped(
|
||||
json!([
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": [{"type": "thinking", "thinking": "plan", "signature": "sig"}]}
|
||||
]),
|
||||
json!([{"role": "user", "content": "hi"}])
|
||||
)]
|
||||
#[case::history_without_thinking_is_kept(
|
||||
json!([{"role": "assistant", "content": [{"type": "text", "text": "hello"}]}]),
|
||||
json!([{"role": "assistant", "content": [{"type": "text", "text": "hello"}]}])
|
||||
)]
|
||||
fn strip_thinking_blocks_rewrites(#[case] input: Value, #[case] expected: Value) {
|
||||
assert_eq!(apply(strip_thinking_blocks, input), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::with_results(json!([{"type": "web_search_result", "url": "u", "title": "Rome", "snippet": "s", "page_age": null}]))]
|
||||
#[case::without_results(json!([]))]
|
||||
|
|
|
|||
|
|
@ -4,7 +4,10 @@ use litellm_types::llms::anthropic_messages::{
|
|||
};
|
||||
|
||||
use crate::{
|
||||
anthropic::experimental_pass_through::messages::thinking::ThinkingContext,
|
||||
anthropic::{
|
||||
common_utils::{is_anthropic_invalid_thinking_block_error, strip_thinking_blocks},
|
||||
experimental_pass_through::messages::thinking::ThinkingContext,
|
||||
},
|
||||
base_llm::chat::transformation::Error,
|
||||
};
|
||||
|
||||
|
|
@ -103,6 +106,21 @@ pub trait BaseAnthropicMessagesConfig: Sync {
|
|||
fn request_headers(&self, headers: Headers, _request: &AnthropicMessagesRequest) -> Headers {
|
||||
headers
|
||||
}
|
||||
|
||||
fn request_after_http_error(
|
||||
&self,
|
||||
status: u16,
|
||||
error_body: &str,
|
||||
request: AnthropicMessagesRequest,
|
||||
) -> Option<AnthropicMessagesRequest> {
|
||||
(status == 400 && is_anthropic_invalid_thinking_block_error(error_body)).then(|| {
|
||||
AnthropicMessagesRequest {
|
||||
thinking: None,
|
||||
messages: strip_thinking_blocks(request.messages),
|
||||
..request
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
|
@ -292,4 +310,58 @@ mod tests {
|
|||
};
|
||||
assert_eq!(config.authenticate(forwarded, api_key, &|_| None), expected);
|
||||
}
|
||||
|
||||
const INVALID_SIGNATURE: &str = r#"{"type":"error","error":{"type":"invalid_request_error","message":"messages.1.content.0: Invalid `signature` in `thinking` block"}}"#;
|
||||
|
||||
fn replayed_thinking_request() -> AnthropicMessagesRequest {
|
||||
serde_json::from_value(serde_json::json!({
|
||||
"model": "claude",
|
||||
"max_tokens": 2048,
|
||||
"thinking": {"type": "enabled", "budget_tokens": 1024},
|
||||
"output_config": {"effort": "high"},
|
||||
"messages": [
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": [
|
||||
{"type": "thinking", "thinking": "plan", "signature": "stale"},
|
||||
{"type": "text", "text": "hello"}
|
||||
]},
|
||||
{"role": "user", "content": "again"}
|
||||
]
|
||||
}))
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_thinking_signature_is_retried_without_thinking() {
|
||||
let retried = DefaultsConfig
|
||||
.request_after_http_error(400, INVALID_SIGNATURE, replayed_thinking_request())
|
||||
.map(|request| serde_json::to_value(request).unwrap());
|
||||
assert_eq!(
|
||||
retried,
|
||||
Some(serde_json::json!({
|
||||
"model": "claude",
|
||||
"max_tokens": 2048,
|
||||
"output_config": {"effort": "high"},
|
||||
"messages": [
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": "hello"}]},
|
||||
{"role": "user", "content": "again"}
|
||||
]
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::server_error_with_the_same_text(500, INVALID_SIGNATURE)]
|
||||
#[case::unrelated_bad_request(400, "rate limit exceeded")]
|
||||
fn other_http_errors_are_not_retried(#[case] status: u16, #[case] error_body: &str) {
|
||||
assert_eq!(
|
||||
DefaultsConfig.request_after_http_error(
|
||||
status,
|
||||
error_body,
|
||||
replayed_thinking_request()
|
||||
),
|
||||
None
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -204,3 +204,109 @@ def test_native_sync_messages_returns_the_provider_message(messages_server: Reco
|
|||
assert_served_natively(messages_server)
|
||||
assert response["content"] == MESSAGES_RESPONSE["content"]
|
||||
assert len(recorder.wait_for("log_success_event")) == 1
|
||||
|
||||
|
||||
INVALID_SIGNATURE: Final = ResponseSpec(
|
||||
body={
|
||||
"type": "error",
|
||||
"error": {
|
||||
"type": "invalid_request_error",
|
||||
"message": "messages.1.content.0: Invalid `signature` in `thinking` block",
|
||||
},
|
||||
},
|
||||
status=400,
|
||||
)
|
||||
REPLAYED_THINKING: Final = (
|
||||
{"role": "user", "content": "hi"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "thinking", "thinking": "plan", "signature": "from-another-deployment"},
|
||||
{"type": "text", "text": "hello"},
|
||||
],
|
||||
},
|
||||
{"role": "user", "content": "again"},
|
||||
)
|
||||
WITHOUT_THINKING: Final = [
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": "hello"}]},
|
||||
{"role": "user", "content": "again"},
|
||||
]
|
||||
|
||||
|
||||
def replayed_thinking(server: RecordingServer, **kwargs: object) -> dict[str, object]:
|
||||
return arguments(
|
||||
server,
|
||||
model="anthropic/claude-haiku-4-5",
|
||||
messages=[dict(message) for message in REPLAYED_THINKING],
|
||||
max_tokens=2048,
|
||||
thinking={"type": "enabled", "budget_tokens": 1024},
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalid_thinking_signature_is_resent_once_without_thinking_and_logged_as_resent(
|
||||
messages_server: RecordingServer,
|
||||
) -> None:
|
||||
messages_server.expected_requests = 2
|
||||
messages_server.enqueue(INVALID_SIGNATURE)
|
||||
recorder: Final = RecordingLogger()
|
||||
|
||||
response: Final = await litellm.anthropic.messages.acreate(
|
||||
**replayed_thinking(messages_server, callbacks=[recorder])
|
||||
)
|
||||
|
||||
assert response["content"] == MESSAGES_RESPONSE["content"]
|
||||
first, second = (request.body for request in messages_server.requests)
|
||||
assert first["thinking"] == {"type": "enabled", "budget_tokens": 1024}
|
||||
assert first["messages"] == list(REPLAYED_THINKING)
|
||||
assert second == {key: value for key, value in first.items() if key not in {"thinking", "messages"}} | {
|
||||
"messages": WITHOUT_THINKING
|
||||
}
|
||||
success: Final = await recorder.wait_for_async("async_log_success_event")
|
||||
assert len(success) == 1
|
||||
assert success[0].kwargs["additional_args"]["complete_input_dict"] == second
|
||||
assert "log_failure_event" not in recorder.names
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalid_thinking_signature_twice_surfaces_the_second_error(messages_server: RecordingServer) -> None:
|
||||
messages_server.expected_requests = 2
|
||||
messages_server.enqueue(INVALID_SIGNATURE)
|
||||
messages_server.enqueue(
|
||||
ResponseSpec(
|
||||
body={"type": "error", "error": {"type": "invalid_request_error", "message": "second"}}, status=400
|
||||
)
|
||||
)
|
||||
|
||||
with pytest.raises(litellm.BadRequestError, match="second"):
|
||||
await litellm.anthropic.messages.acreate(**replayed_thinking(messages_server))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_other_bad_requests_are_not_resent(messages_server: RecordingServer) -> None:
|
||||
messages_server.enqueue(
|
||||
ResponseSpec(
|
||||
body={"type": "error", "error": {"type": "invalid_request_error", "message": "prompt is too long"}},
|
||||
status=400,
|
||||
)
|
||||
)
|
||||
|
||||
with pytest.raises(litellm.BadRequestError, match="prompt is too long"):
|
||||
await litellm.anthropic.messages.acreate(**replayed_thinking(messages_server))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalid_thinking_signature_on_a_stream_is_resent_and_relayed(messages_server: RecordingServer) -> None:
|
||||
messages_server.expected_requests = 2
|
||||
messages_server.enqueue(INVALID_SIGNATURE)
|
||||
messages_server.enqueue(STREAM)
|
||||
|
||||
stream: Final = await litellm.anthropic.messages.acreate(**replayed_thinking(messages_server, stream=True))
|
||||
assert isinstance(stream, AsyncIterator)
|
||||
chunks: Final = [chunk async for chunk in stream]
|
||||
|
||||
assert b"".join(chunks) == sse_payload()
|
||||
assert messages_server.requests[1].body["messages"] == WITHOUT_THINKING
|
||||
assert messages_server.requests[1].body["stream"] is True
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue