Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
267 changes: 267 additions & 0 deletions src/libraries/rust/stargate/crates/pylon-lib/src/bringup.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1245,6 +1245,273 @@ mod tests {
progress.await.unwrap();
}

/// Character text whose estimate is one token per character.
fn estimated_tokens(count: usize) -> String {
"\u{e9}".repeat(count)
}

fn reasoning_event(text: &str) -> Value {
serde_json::json!({
"object": "chat.completion.chunk",
"choices": [{"index": 0, "delta": {"reasoning_content": text}}],
})
}

fn usage_event(text: &str, completion_tokens: u32) -> Value {
serde_json::json!({
"object": "chat.completion.chunk",
"choices": [{"index": 0, "delta": {"reasoning_content": text}}],
"usage": {"completion_tokens": completion_tokens},
})
}

#[derive(Clone, Copy, PartialEq, Eq)]
enum CanaryEnd {
Done,
Eof,
Stall,
}

/// Serves the given canary events, then ends the stream as `end` says.
async fn spawn_canary_event_server(events: Vec<Value>, end: CanaryEnd) -> TestHttpServer {
use futures::StreamExt as _;

TestHttpServer::spawn(Router::new().route(
"/v1/chat/completions",
post(move || async move {
let mut body = events
.iter()
.map(|event| format!("data: {event}\n\n"))
.collect::<String>();
if end == CanaryEnd::Done {
body.push_str("data: [DONE]\n\n");
}
let first = futures::stream::once(async move {
Ok::<_, std::io::Error>(bytes::Bytes::from(body))
});
let body = if end == CanaryEnd::Stall {
axum::body::Body::from_stream(first.chain(futures::stream::pending()))
} else {
axum::body::Body::from_stream(first)
};
([("content-type", "text/event-stream")], body)
}),
))
.await
}

async fn send_test_canary(events: Vec<Value>, end: CanaryEnd) -> Result<(), BringupError> {
let server = spawn_canary_event_server(events, end).await;
let result = send_canary_request(
&reqwest::Client::new(),
server.as_str(),
&test_generation(),
Duration::from_millis(200),
7,
)
.await;
server.shutdown().await;
result
}

#[tokio::test]
async fn canary_exact_usage_at_threshold_settles_a_split_estimate_overshoot() {
send_test_canary(
vec![reasoning_event(&estimated_tokens(15)), usage_event("", 7)],
CanaryEnd::Done,
)
.await
.expect("exact usage at the cap should override a higher estimate");
}

#[tokio::test]
async fn canary_exact_usage_at_threshold_settles_an_overshoot_in_the_same_event() {
send_test_canary(vec![usage_event(&estimated_tokens(15), 7)], CanaryEnd::Done)
.await
.expect("exact usage at the cap should override a higher estimate");
}

#[tokio::test]
async fn canary_mid_stream_usage_then_final_usage_at_threshold_passes() {
send_test_canary(
vec![
usage_event(&estimated_tokens(5), 5),
reasoning_event(&estimated_tokens(4)),
usage_event("", 7),
],
CanaryEnd::Done,
)
.await
.expect("final exact usage at the cap should settle the estimated tail");
}

#[tokio::test]
async fn canary_judges_a_completed_stream_by_exact_usage_despite_a_later_estimated_tail() {
send_test_canary(
vec![
usage_event(&estimated_tokens(7), 7),
reasoning_event(&estimated_tokens(9)),
],
CanaryEnd::Done,
)
.await
.expect("accepted exact usage should decide the verdict");
}

#[tokio::test]
async fn canary_ignores_regressed_exact_usage() {
send_test_canary(
vec![usage_event(&estimated_tokens(5), 5), usage_event("", 0)],
CanaryEnd::Done,
)
.await
.expect("regressed usage should keep the prior exact count");
}

#[tokio::test]
async fn canary_failure_event_with_usage_above_threshold_is_not_runaway() {
let error = send_test_canary(
vec![serde_json::json!({
"type": "error",
"usage": {"completion_tokens": 8},
})],
CanaryEnd::Done,
)
.await
.expect_err("a failed stream should fail");

assert!(
matches!(error, BringupError::InvalidResponse(_)),
"unexpected canary error: {error}"
);
}

#[tokio::test]
async fn canary_exact_usage_above_threshold_is_runaway() {
let error = send_test_canary(
vec![reasoning_event(&estimated_tokens(3)), usage_event("", 8)],
CanaryEnd::Done,
)
.await
.expect_err("exact usage above the cap should fail");

assert!(
matches!(error, BringupError::RunawayGeneration { tokens: 8 }),
"unexpected canary error: {error}"
);
}

#[tokio::test]
async fn canary_exact_usage_above_threshold_fails_before_the_stream_ends() {
let error = send_test_canary(vec![usage_event("", 8)], CanaryEnd::Stall)
.await
.expect_err("exact usage above the cap should fail immediately");

assert!(
matches!(error, BringupError::RunawayGeneration { tokens: 8 }),
"unexpected canary error: {error}"
);
}

#[tokio::test]
async fn canary_without_usage_uses_the_estimate_at_completion() {
send_test_canary(vec![reasoning_event(&estimated_tokens(7))], CanaryEnd::Done)
.await
.expect("an estimate at the cap should pass");

let error = send_test_canary(vec![reasoning_event(&estimated_tokens(8))], CanaryEnd::Done)
.await
.expect_err("an unsettled estimate above the cap should fail");

assert!(
matches!(error, BringupError::RunawayGeneration { tokens: 8 }),
"unexpected canary error: {error}"
);
}

#[tokio::test]
async fn canary_stall_after_estimate_overshoot_is_a_timeout() {
let error = send_test_canary(
vec![reasoning_event(&estimated_tokens(9))],
CanaryEnd::Stall,
)
.await
.expect_err("a stalled canary should fail");

assert!(error.is_timeout(), "unexpected canary error: {error}");
}

#[tokio::test]
async fn canary_eof_after_estimate_overshoot_is_an_invalid_response() {
let error = send_test_canary(vec![reasoning_event(&estimated_tokens(9))], CanaryEnd::Eof)
.await
.expect_err("a stream without [DONE] should fail");

assert!(
matches!(error, BringupError::InvalidResponse(_)),
"unexpected canary error: {error}"
);
}

#[tokio::test]
async fn canary_completion_event_without_done_is_an_invalid_response() {
let error = send_test_canary(
vec![
reasoning_event("2"),
serde_json::json!({"type": "response.completed"}),
],
CanaryEnd::Eof,
)
.await
.expect_err("a canary stream must end with [DONE]");

assert!(
matches!(error, BringupError::InvalidResponse(_)),
"unexpected canary error: {error}"
);
}

#[tokio::test]
async fn canary_read_error_after_estimate_overshoot_is_not_runaway() {
use futures::StreamExt as _;

let server = TestHttpServer::spawn(Router::new().route(
"/v1/chat/completions",
post(|| async {
let event = format!("data: {}\n\n", reasoning_event(&estimated_tokens(9)));
let first = futures::stream::once(async move {
Ok::<_, std::io::Error>(bytes::Bytes::from(event))
});
let reset = futures::stream::once(async {
tokio::time::sleep(Duration::from_millis(50)).await;
Err(std::io::Error::other("connection reset"))
});
let body = first.chain(reset);
(
[("content-type", "text/event-stream")],
axum::body::Body::from_stream(body),
)
}),
))
.await;

let error = send_canary_request(
&reqwest::Client::new(),
server.as_str(),
&test_generation(),
Duration::from_secs(1),
7,
)
.await
.expect_err("a broken canary stream should fail");

assert!(
matches!(error, BringupError::InvalidResponse(_)),
"unexpected canary error: {error}"
);
server.shutdown().await;
}

#[tokio::test]
async fn canary_stream_stall_is_classified_as_a_timeout() {
use futures::StreamExt as _;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -165,13 +165,18 @@ pub(super) async fn send_canary_request(
DEFAULT_MAX_SSE_BUFFER_BYTES,
false,
);
let threshold = u64::from(canary_max_generation_threshold);
let mut output_tokens = OutputTokenParser::new();
let mut observed_tokens = 0_u64;
let mut estimated_tokens = 0_u64;
let mut exact_tokens = None;
let runaway = |tokens: u64| BringupError::RunawayGeneration {
tokens: u32::try_from(tokens).unwrap_or(u32::MAX),
};
let mut completed = false;
while let Some(message) = messages.next().await {
let message = message.map_err(canary_stream_error)?;
if let Some(generated_output) = message.facts.generated_output {
observed_tokens = output_tokens
estimated_tokens = output_tokens
.observe_generated_characters(generated_output.characters)
.displayed_tokens;
}
Expand All @@ -181,27 +186,33 @@ pub(super) async fn send_canary_request(
.and_then(|usage| usage.output_tokens)
&& output_tokens.observe_exact_output_tokens(tokens) == ExactOutputUpdate::Applied
{
observed_tokens = tokens;
}
if observed_tokens > u64::from(canary_max_generation_threshold) {
return Err(BringupError::RunawayGeneration {
tokens: u32::try_from(observed_tokens).unwrap_or(u32::MAX),
});
exact_tokens = Some(tokens);
}
match message.facts.terminal {
Some(RelayOutcome::Complete) => {
Some(RelayOutcome::Complete) if message.facts.done_sentinel => {
completed = true;
break;
}
Some(RelayOutcome::Failed) => break,
None => {}
Some(RelayOutcome::Complete) | None => {}
}
if let Some(tokens) = exact_tokens
&& tokens > threshold
{
return Err(runaway(tokens));
}
}
// A character estimate can overshoot a response that stopped at the cap,
// so it is judged only when a completed stream reports no exact usage.
let observed_tokens = exact_tokens.unwrap_or(estimated_tokens);
if !completed || observed_tokens == 0 {
return Err(BringupError::InvalidResponse(
"canary stream must contain output and end with [DONE]".to_string(),
));
}
if observed_tokens > threshold {
return Err(runaway(observed_tokens));
}
Ok(())
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,8 @@ pub(crate) struct SseEventFacts {
pub(crate) reasoning_output_observed: bool,
pub(crate) calibration_ineligible: bool,
pub(crate) terminal: Option<RelayOutcome>,
/// The event's data is the Chat Completions `[DONE]` sentinel.
pub(crate) done_sentinel: bool,
}

#[derive(Debug, PartialEq)]
Expand Down Expand Up @@ -346,6 +348,7 @@ fn classify_sse_event(
calibration_ineligible: output_token_calibration_enabled
&& event_name.is_some_and(|event_name| event_name.starts_with("response.")),
terminal: Some(terminal),
done_sentinel: true,
..SseEventFacts::default()
},
);
Expand Down Expand Up @@ -394,6 +397,7 @@ fn classify_sse_event(
terminal_outcome(json_event_type),
terminal_outcome(event_name),
),
done_sentinel: false,
};
(parsed, facts)
}
Expand Down
2 changes: 1 addition & 1 deletion src/libraries/rust/stargate/crates/pylon/src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -107,7 +107,7 @@ struct Args {
/// Interval between active canary requests in milliseconds. Models with request progress within the interval skip the canary. `0` disables active canaries
#[arg(long, default_value_t = 5000, value_name = "MS")]
active_canary_interval_ms: u64,
/// Treat canary responses that generate this many tokens as runaway generation
/// Canary `max_tokens`. Generating more tokens than this is runaway generation; exact usage overrides the output estimate
#[arg(long, default_value_t = 237, value_name = "TOKENS")]
canary_max_generation_threshold: u32,
/// Initial calibration request count; doubles after each completed load step
Expand Down
Loading