use serde::de::DeserializeOwned;
use std::time::Instant;
use crate::prompt::load_prompt;
use crate::retry::{
FailureClass, RetryExhausted, RetryFailureRecord, RetryLoop, RetryPolicy, fail_exhausted,
};
use crate::util::json::parse_fenced_json;
use crate::{ChatMessage, ChatRequest, ExtractionValidator};
#[expect(clippy::too_many_lines)]
pub(crate) async fn retry_extract_structured_scoped<T: DeserializeOwned>(
history: &[ChatMessage],
extraction_prompt: &str,
params: &ChatRequest,
validate: Option<&ExtractionValidator<T>>,
policy_override: Option<&RetryPolicy>,
) -> Result<T, RetryExhausted> {
let mut extraction_history = history.to_vec();
let operation_started = Instant::now();
if !extraction_prompt.is_empty() {
extraction_history.push(ChatMessage::user(extraction_prompt));
}
let record_request = ChatRequest {
messages: vec![],
..params.clone()
};
let retry_prompt = load_prompt("extraction/retry.md");
let policy = policy_override
.cloned()
.unwrap_or_else(RetryPolicy::current);
let mut loop_state = RetryLoop::new(&policy);
let mut last_raw: Option<String> = None;
for attempt in 1..=policy.max_attempts {
if loop_state.expired() {
let exhausted = RetryExhausted::with_last_raw(
loop_state.into_failures(),
FailureClass::WallClockExceeded,
last_raw,
);
return fail_exhausted(&record_request, operation_started, exhausted).await;
}
let request = ChatRequest {
messages: extraction_history.clone(),
..record_request.clone()
};
match crate::providers::chat_scoped(request, policy.idle_timeout, loop_state.deadline())
.await
{
Ok(response) => {
last_raw = Some(response.text_or_empty().to_string());
let (class, detail): (FailureClass, String) = if response.tool_calls.is_empty() {
let raw = last_raw.as_deref().unwrap_or_default();
match parse_fenced_json::<T>(raw) {
Ok(result) => {
if let Some(validate) = validate
&& let Err(msg) = validate(&result)
{
(
FailureClass::OutOfRangeScore,
format!("extracted value rejected by validation: {msg}"),
)
} else {
crate::stats::record_llm_success(
&record_request,
operation_started,
attempt,
&response,
)
.await;
return Ok(result);
}
}
Err(e) => (
FailureClass::Parse,
format!("failed to parse extracted JSON: {e}"),
),
}
} else {
(
FailureClass::Parse,
"extraction attempt returned a tool call instead of JSON".to_string(),
)
};
let err = anyhow::anyhow!("{detail}");
let rec = RetryFailureRecord::new_simple(class, &err, None);
loop_state.record(rec);
if attempt < policy.max_attempts {
extraction_history
.push(ChatMessage::assistant(last_raw.clone().unwrap_or_default()));
extraction_history.push(ChatMessage::user(retry_prompt.as_str()));
}
}
Err(scoped_err) => {
last_raw = None;
let non_retryable = !scoped_err.class.is_retryable();
loop_state.record(scoped_err.record);
if non_retryable {
let exhausted = RetryExhausted::with_last_raw(
loop_state.into_failures(),
scoped_err.class,
last_raw,
);
return fail_exhausted(&record_request, operation_started, exhausted).await;
}
}
}
if let Err(class) = loop_state.sleep_between(attempt).await {
let exhausted =
RetryExhausted::with_last_raw(loop_state.into_failures(), class, last_raw);
return fail_exhausted(&record_request, operation_started, exhausted).await;
}
}
let final_class = loop_state.final_class();
let exhausted =
RetryExhausted::with_last_raw(loop_state.into_failures(), final_class, last_raw);
fail_exhausted(&record_request, operation_started, exhausted).await
}
#[cfg(test)]
mod tests {
use super::*;
use crate::retry::tiny_test_policy;
use crate::util::test::{FakeProvider, install_fake_provider, retry_tests_lock};
use crate::{ChatMessage, ChatRequest};
use std::sync::Arc;
#[derive(serde::Deserialize, Debug, PartialEq)]
struct FakeVerdict {
score: u8,
}
fn score_validator(v: &FakeVerdict) -> Result<(), String> {
if v.score <= 10 {
Ok(())
} else {
Err(format!("score {} out of range", v.score))
}
}
fn test_params() -> ChatRequest {
ChatRequest {
messages: vec![],
tools: None,
model: "test-model".to_string(),
max_tokens: None,
reasoning_effort: None,
provider_order: None,
meta: None,
}
}
fn history() -> Vec<ChatMessage> {
vec![ChatMessage::user("analyze the ticket")]
}
async fn extract_with(
fake: Arc<FakeProvider>,
) -> Result<FakeVerdict, crate::retry::RetryExhausted> {
let provider: Arc<dyn crate::Provider> = fake.clone();
let _guard = install_fake_provider(provider);
retry_extract_structured_scoped::<FakeVerdict>(
&history(),
"return JSON verdict",
&test_params(),
Some(&score_validator),
None,
)
.await
}
#[tokio::test]
#[expect(clippy::await_holding_lock)] async fn recovers_after_consecutive_transport_errors() {
let _guard = retry_tests_lock();
let _policy_guard = crate::util::test::install_test_retry_policy(tiny_test_policy());
let fake = Arc::new(
FakeProvider::new()
.err(crate::retry::FailureClass::Transport, "transport error 1")
.err(crate::retry::FailureClass::Transport, "transport error 2")
.ok(r#"{"score": 8}"#),
);
let result = extract_with(fake).await;
let verdict = result.expect("should recover after 2 transport errors");
assert_eq!(verdict.score, 8);
}
#[tokio::test]
#[expect(clippy::await_holding_lock)] async fn recovers_after_truncated_envelope_parse_errors() {
let _guard = retry_tests_lock();
let _policy_guard = crate::util::test::install_test_retry_policy(tiny_test_policy());
let fake = Arc::new(
FakeProvider::new()
.err(
crate::retry::FailureClass::TruncatedEnvelope,
"EOF while parsing a value at line 317",
)
.ok(r#"{"score": 9}"#),
);
let result = extract_with(fake).await;
assert_eq!(result.expect("recovery").score, 9);
}
#[tokio::test]
#[expect(clippy::await_holding_lock)] async fn recovers_after_llm_parse_failure_via_reprompt() {
let _guard = retry_tests_lock();
let _policy_guard = crate::util::test::install_test_retry_policy(tiny_test_policy());
let fake = Arc::new(
FakeProvider::new()
.ok("this is not JSON at all")
.ok(r#"{"score": 7}"#),
);
let result = extract_with(fake).await;
assert_eq!(result.expect("recovery after re-prompt").score, 7);
}
#[tokio::test]
#[expect(clippy::await_holding_lock)] async fn exhausts_attempts_with_bounded_call_count() {
let _guard = retry_tests_lock();
let _policy_guard = crate::util::test::install_test_retry_policy(tiny_test_policy()); let fake = Arc::new(
FakeProvider::new()
.err(crate::retry::FailureClass::Transport, "always down")
.err(crate::retry::FailureClass::Transport, "always down")
.err(crate::retry::FailureClass::Transport, "always down"),
);
let result = extract_with(fake).await;
let failure = result.expect_err("must exhaust");
assert_eq!(failure.final_class, crate::retry::FailureClass::Transport);
assert_eq!(failure.failures.len(), 3);
assert_eq!(failure.last_raw, None, "no text ever produced");
}
#[tokio::test]
#[expect(clippy::await_holding_lock)] async fn non_retryable_error_propagates_immediately() {
let _guard = retry_tests_lock();
let _policy_guard = crate::util::test::install_test_retry_policy(tiny_test_policy());
let fake = Arc::new(FakeProvider::new().err(
crate::retry::FailureClass::NonRetryable,
"insufficient balance",
));
let result = extract_with(fake).await;
let failure = result.expect_err("non-retryable must abort");
assert_eq!(
failure.final_class,
crate::retry::FailureClass::NonRetryable
);
assert_eq!(failure.failures.len(), 1, "single call, no retries");
}
#[tokio::test]
#[expect(clippy::await_holding_lock)] async fn transport_final_failure_clears_earlier_text_from_last_raw() {
let _guard = retry_tests_lock();
let _policy_guard = crate::util::test::install_test_retry_policy(tiny_test_policy()); let fake = Arc::new(
FakeProvider::new()
.ok("this is not JSON")
.err(crate::retry::FailureClass::Transport, "body read failed")
.err(crate::retry::FailureClass::Transport, "body read failed"),
);
let result = extract_with(fake).await;
let failure = result.expect_err("must exhaust on transport");
assert_eq!(failure.final_class, crate::retry::FailureClass::Transport);
assert_eq!(
failure.last_raw, None,
"earlier attempt text must not be presented as the last attempt"
);
assert_eq!(failure.failures.len(), 3);
assert_eq!(failure.failures[0].class, crate::retry::FailureClass::Parse);
assert_eq!(
failure.failures[1].class,
crate::retry::FailureClass::Transport
);
assert_eq!(
failure.failures[2].class,
crate::retry::FailureClass::Transport
);
}
#[tokio::test]
#[expect(clippy::await_holding_lock)] async fn tool_call_final_attempt_keeps_empty_last_raw() {
let _guard = retry_tests_lock();
let _policy_guard = crate::util::test::install_test_retry_policy(tiny_test_policy()); let fake = Arc::new(
FakeProvider::new().ok("not json").ok("not json").ok(""), );
let result = extract_with(fake).await;
let failure = result.expect_err("must exhaust");
assert_eq!(failure.last_raw, Some(String::new()));
assert_eq!(
failure.failures.last().map(|r| r.class),
Some(crate::retry::FailureClass::Parse)
);
}
#[tokio::test]
#[expect(clippy::await_holding_lock)] async fn attempts_byte_identical_except_reprompt() {
let _guard = retry_tests_lock();
let _policy_guard = crate::util::test::install_test_retry_policy(tiny_test_policy());
let fake = Arc::new(
FakeProvider::new()
.err(crate::retry::FailureClass::Transport, "transient")
.ok("garbage text")
.ok(r#"{"score": 6}"#),
);
let result = extract_with(fake.clone()).await;
assert_eq!(result.expect("recovers").score, 6);
let fingerprints = fake.request_fingerprints.lock().unwrap().clone();
assert_eq!(fingerprints.len(), 3);
assert_eq!(
fingerprints[0], fingerprints[1],
"transport retry must be byte-identical"
);
assert_ne!(
fingerprints[1], fingerprints[2],
"parse-failure re-prompt must extend the request"
);
}
#[tokio::test]
#[expect(clippy::await_holding_lock)] async fn score_out_of_range_never_passes_and_reprompts() {
let _guard = retry_tests_lock();
let _policy_guard = crate::util::test::install_test_retry_policy(tiny_test_policy());
let fake = Arc::new(
FakeProvider::new()
.ok(r#"{"score": 255}"#) .ok(r#"{"score": 5}"#),
);
let result = extract_with(fake).await;
let verdict = result.expect("re-prompted after out-of-range score");
assert_eq!(verdict.score, 5);
}
#[tokio::test]
#[expect(clippy::await_holding_lock)] async fn score_out_of_range_all_attempts_classifies_as_out_of_range() {
let _guard = retry_tests_lock();
let _policy_guard = crate::util::test::install_test_retry_policy(tiny_test_policy()); let fake = Arc::new(
FakeProvider::new()
.ok(r#"{"score": 101}"#) .ok(r#"{"score": 101}"#)
.ok(r#"{"score": 101}"#),
);
let result = extract_with(fake).await;
let failure = result.expect_err("garbage score must never pass");
assert_eq!(
failure.final_class,
crate::retry::FailureClass::OutOfRangeScore,
"all-attempts-out-of-range must classify explicitly"
);
assert!(
failure
.failures
.iter()
.all(|r| r.class == crate::retry::FailureClass::OutOfRangeScore)
);
assert_eq!(failure.failures.len(), 3);
}
#[tokio::test]
#[expect(clippy::await_holding_lock)] async fn wall_clock_cap_binds_before_attempt_exhaustion() {
let _guard = retry_tests_lock();
let _policy_guard =
crate::util::test::install_test_retry_policy(crate::retry::RetryPolicy {
max_attempts: 7,
base_backoff_ms: 1_000,
max_backoff_ms: 1_000,
operation_timeout: std::time::Duration::from_millis(100),
idle_timeout: std::time::Duration::from_secs(1),
});
let fake = Arc::new(
FakeProvider::new()
.err(crate::retry::FailureClass::Transport, "slow outage")
.err(crate::retry::FailureClass::Transport, "slow outage"),
);
let result = extract_with(fake).await;
let failure = result.expect_err("wall-clock cap must bind");
assert_eq!(
failure.final_class,
crate::retry::FailureClass::WallClockExceeded
);
assert!(
failure.failures.len() <= 2,
"cap must stop the loop before 7 attempts"
);
}
#[tokio::test]
#[expect(clippy::await_holding_lock)] async fn extraction_rows_record_base_params_and_retry_context() {
let _guard = retry_tests_lock();
let _policy_guard = crate::util::test::install_test_retry_policy(tiny_test_policy());
let (store, _tmp) = crate::open_test_store!(crate::logs::LogStore, "log");
let _store_guard = crate::util::test::install_test_log_store(store.clone());
let params = ChatRequest {
meta: Some(crate::ChatRequestMeta {
purpose: "extraction",
agent_id: "extraction-flag-test".to_string(),
role: "reviewer".to_string(),
workspace: "ws1".to_string(),
ticket_id: None,
}),
..test_params()
};
let fake = Arc::new(FakeProvider::new().ok(r#"{"score": 8}"#));
let provider: Arc<dyn crate::Provider> = fake.clone();
let _fake_guard = install_fake_provider(provider);
let result = retry_extract_structured_scoped::<FakeVerdict>(
&history(),
"return JSON verdict",
¶ms,
Some(&score_validator),
None,
)
.await;
assert_eq!(result.expect("extraction succeeds").score, 8);
let fake = Arc::new(FakeProvider::new().ok("garbage").ok(r#"{"score": 7}"#));
let provider: Arc<dyn crate::Provider> = fake.clone();
let _fake_guard = install_fake_provider(provider);
let result = retry_extract_structured_scoped::<FakeVerdict>(
&history(),
"return JSON verdict",
¶ms,
Some(&score_validator),
None,
)
.await;
assert_eq!(result.expect("recovers after re-prompt").score, 7);
let fake = Arc::new(
FakeProvider::new()
.err(crate::retry::FailureClass::Transport, "always down")
.err(crate::retry::FailureClass::Transport, "always down")
.err(crate::retry::FailureClass::Transport, "always down"),
);
let provider: Arc<dyn crate::Provider> = fake.clone();
let _fake_guard = install_fake_provider(provider);
let result = retry_extract_structured_scoped::<FakeVerdict>(
&history(),
"return JSON verdict",
¶ms,
Some(&score_validator),
None,
)
.await;
assert!(
result.is_err(),
"all-attempt failure must not produce a verdict"
);
let rows = store
.conn
.query(
"SELECT cost, upstream_provider, success \
FROM llm_requests WHERE agent_id = ?1 ORDER BY rowid",
crate::turso::params!["extraction-flag-test"],
)
.await
.expect("query recorded rows");
let mut rows = rows.into_iter();
let success = rows.next().expect("success row must exist");
assert_eq!(success.get::<i64>(2).expect("success"), 1);
let recovery = rows.next().expect("recovery row must exist");
assert_eq!(recovery.get::<i64>(2).expect("success"), 1);
let failure = rows.next().expect("failure row must exist");
assert_eq!(failure.get::<i64>(2).expect("success"), 0);
assert_eq!(
failure.get::<Option<f64>>(0).expect("cost"),
None,
"no envelope on failure path — cost must stay NULL"
);
assert_eq!(
failure.get::<Option<String>>(1).expect("upstream_provider"),
None,
"no envelope on failure path — provider must stay NULL"
);
assert!(rows.next().is_none(), "exactly three rows");
}
}