fix: harden concurrency limits and high-RPM runtime paths

Bound request, stream, queue, and shutdown resource lifetimes. Reduce scheduler and Redis hot-path work and isolate database maintenance. Include regression coverage, load probes, and concurrency audit results.
This commit is contained in:
elky
2026-09-10 08:14:58 +08:00
parent 361952ada9
commit ecc16673eb
149 changed files with 27963 additions and 1926 deletions
@@ -24,6 +24,8 @@ struct GeminiProviderToolResultState {
#[derive(Default)]
pub struct GeminiProviderState {
terminal_observation_only: bool,
observed_tool_calls: bool,
response_id: Option<String>,
model: Option<String>,
started: bool,
@@ -37,6 +39,13 @@ pub struct GeminiProviderState {
}
impl GeminiProviderState {
pub(crate) fn terminal_observation() -> Self {
Self {
terminal_observation_only: true,
..Self::default()
}
}
fn identity(&self, report_context: &Value) -> (String, String) {
resolve_identity(
self.response_id.as_deref(),
@@ -133,14 +142,21 @@ impl GeminiProviderState {
let Some(part_object) = part.as_object() else {
continue;
};
let reasoning_signature = part_object
.get("thoughtSignature")
.or_else(|| part_object.get("thought_signature"))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let reasoning_signature = if self.terminal_observation_only {
None
} else {
part_object
.get("thoughtSignature")
.or_else(|| part_object.get("thought_signature"))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
};
if let Some(text) = render_gemini_part_as_text(part_object) {
if self.terminal_observation_only {
continue;
}
let is_reasoning = part_object
.get("thought")
.and_then(Value::as_bool)
@@ -193,6 +209,9 @@ impl GeminiProviderState {
.or_else(|| part_object.get("function_response"))
.and_then(Value::as_object)
{
if self.terminal_observation_only {
continue;
}
let tool_use_id = function_response
.get("id")
.and_then(Value::as_str)
@@ -240,6 +259,9 @@ impl GeminiProviderState {
else {
if let Some(content_part) = canonical_content_part_from_gemini_part(part_object)
{
if self.terminal_observation_only {
continue;
}
let should_emit = self
.content_parts
.get(&index)
@@ -260,6 +282,10 @@ impl GeminiProviderState {
}
continue;
};
if self.terminal_observation_only {
self.observed_tool_calls = true;
continue;
}
let tool_state = self.tool_calls.entry(index).or_default();
tool_state.call_id = function_call
.get("id")
@@ -329,7 +355,7 @@ impl GeminiProviderState {
if let Some(finish_reason) =
candidate_object.get("finishReason").and_then(Value::as_str)
{
let has_tool_calls = !self.tool_calls.is_empty();
let has_tool_calls = self.observed_tool_calls || !self.tool_calls.is_empty();
let mut finish_reason =
normalize_openai_finish_reason(map_gemini_stream_finish_reason(finish_reason));
if has_tool_calls && finish_reason.as_deref().is_none_or(|value| value == "stop") {
@@ -936,6 +962,232 @@ mod tests {
format!("data: {}\n", value).into_bytes()
}
fn terminal_frames(frames: Vec<CanonicalStreamFrame>) -> Vec<CanonicalStreamFrame> {
frames
.into_iter()
.filter(|frame| {
matches!(
frame.event,
CanonicalStreamEvent::Start
| CanonicalStreamEvent::Finish { .. }
| CanonicalStreamEvent::UnknownEvent(_)
)
})
.collect()
}
fn observation_record(parts: Vec<Value>, finish_reason: Option<&str>) -> Value {
let mut record = json!({
"responseId": "resp_observation",
"modelVersion": "gemini-2.5-pro",
"candidates": [{"content": {"parts": parts}}],
"usageMetadata": {
"promptTokenCount": 22,
"cachedContentTokenCount": 7,
"candidatesTokenCount": 13,
"thoughtsTokenCount": 5,
"totalTokenCount": 40
}
});
if let Some(finish_reason) = finish_reason {
record["candidates"][0]["finishReason"] = json!(finish_reason);
}
record
}
fn assert_terminal_record_matches(
normal: &mut GeminiProviderState,
observer: &mut GeminiProviderState,
record: Value,
) -> Vec<CanonicalStreamFrame> {
let context = json!({"mapped_model": "fallback-model"});
let expected = terminal_frames(
normal
.push_line(&context, data_line(record.clone()))
.expect("normal provider parser"),
);
let actual = observer
.push_line(&context, data_line(record))
.expect("terminal provider parser");
assert_eq!(
actual, expected,
"terminal frames must retain their full payloads"
);
actual
}
fn assert_observer_has_no_content_buffers(observer: &GeminiProviderState) {
assert!(observer.text_parts.is_empty());
assert!(observer.reasoning_parts.is_empty());
assert!(observer.reasoning_signatures.is_empty());
assert!(observer.content_parts.is_empty());
assert!(observer.tool_calls.is_empty());
assert!(observer.tool_results.is_empty());
}
#[test]
fn gemini_terminal_observation_does_not_retain_long_stream_content() {
let mut normal = GeminiProviderState::default();
let mut observer = GeminiProviderState::terminal_observation();
let media = "YQ==".repeat(4_096);
for index in 0..32 {
let text = format!("record-{index}:{}", "x".repeat(16_384));
assert_terminal_record_matches(
&mut normal,
&mut observer,
observation_record(
vec![
json!({"text": text}),
json!({"text": text, "thought": true, "thoughtSignature": media}),
json!({"functionCall": {"id": format!("call-{index}"), "name": "lookup", "args": {"value": text}}}),
json!({"functionResponse": {"name": "lookup", "response": {"result": text}}}),
json!({"inlineData": {"mimeType": "image/png", "data": media}}),
json!({"inline_data": {"mime_type": "audio/wav", "data": media}}),
json!({"inlineData": {"mimeType": "application/pdf", "data": media}}),
],
None,
),
);
assert_observer_has_no_content_buffers(&observer);
assert!(observer.observed_tool_calls);
}
assert!(!normal.text_parts.is_empty());
assert!(!normal.reasoning_parts.is_empty());
assert!(!normal.reasoning_signatures.is_empty());
assert!(!normal.content_parts.is_empty());
assert!(!normal.tool_calls.is_empty());
assert!(normal
.tool_results
.values()
.any(|state| state.content.len() > 16_384));
let frames = assert_terminal_record_matches(
&mut normal,
&mut observer,
observation_record(Vec::new(), Some("STOP")),
);
assert!(frames.iter().any(|frame| matches!(
&frame.event,
CanonicalStreamEvent::Finish { finish_reason: Some(reason), usage: Some(usage) }
if reason == "tool_calls" && usage.input_tokens == 22 && usage.output_tokens == 18
&& usage.cache_read_tokens == 7 && usage.reasoning_tokens == 5 && usage.total_tokens == 40
)));
assert_observer_has_no_content_buffers(&observer);
}
#[test]
fn gemini_terminal_observation_matches_known_and_unknown_part_classification() {
let parts = vec![
Value::Null,
json!({}),
json!({"text": ""}),
json!({"text": 17}),
json!({"text": "", "thought_signature": "signature"}),
json!({"thoughtSignature": "signature"}),
json!({"executableCode": {"language": "python", "code": "print(1)"}}),
json!({"codeExecutionResult": {"output": "1"}}),
json!({"functionResponse": {}}),
json!({"function_response": {"response": ["result"]}}),
json!({"functionCall": {}}),
json!({"functionCall": null}),
json!({"inlineData": {"mimeType": "image/png", "data": "YQ=="}}),
json!({"inline_data": {"mime_type": "audio/wav", "data": "YQ=="}}),
json!({"inlineData": {"mimeType": "application/pdf", "data": "YQ=="}}),
json!({"inlineData": {"mimeType": "", "data": "YQ=="}}),
json!({"inlineData": {"mimeType": "image/png", "data": ""}}),
json!({"file_data": {"file_uri": "gs://test/file", "mime_type": "application/pdf"}}),
json!({"fileData": {"fileUri": ""}}),
json!({"futurePart": {"kept": true}}),
];
for part in parts {
let mut normal = GeminiProviderState::default();
let mut observer = GeminiProviderState::terminal_observation();
assert_terminal_record_matches(
&mut normal,
&mut observer,
json!({"responseId": "outer", "response": observation_record(vec![part], Some("STOP"))}),
);
assert_eq!(observer.response_id.as_deref(), Some("resp_observation"));
assert_observer_has_no_content_buffers(&observer);
}
}
#[test]
fn gemini_terminal_observation_preserves_finish_reasons_and_eof() {
for has_tool in [false, true] {
for reason in [
None,
Some("STOP"),
Some("MAX_TOKENS"),
Some("RECITATION"),
Some("FUTURE_REASON"),
] {
let mut normal = GeminiProviderState::default();
let mut observer = GeminiProviderState::terminal_observation();
let part = if has_tool {
json!({"functionCall": {"name": "lookup", "args": {"query": "test"}}})
} else {
json!({"text": "partial output"})
};
assert_terminal_record_matches(
&mut normal,
&mut observer,
observation_record(vec![part], None),
);
if let Some(reason) = reason {
assert_terminal_record_matches(
&mut normal,
&mut observer,
observation_record(Vec::new(), Some(reason)),
);
}
let context = json!({});
assert_eq!(
observer.finish(&context).expect("observer EOF"),
terminal_frames(normal.finish(&context).expect("normal EOF"))
);
assert!(observer
.finish(&context)
.expect("idempotent EOF")
.is_empty());
assert_observer_has_no_content_buffers(&observer);
}
}
}
#[test]
fn gemini_terminal_observation_preserves_failure_payloads_with_usage() {
for reason in [
"MALFORMED_FUNCTION_CALL",
"UNEXPECTED_TOOL_CALL",
"TOO_MANY_TOOL_CALLS",
"MISSING_THOUGHT_SIGNATURE",
"MALFORMED_RESPONSE",
] {
for content in [
Value::Null,
json!({}),
json!({"parts": []}),
json!({"parts": [{"text": ""}]}),
] {
let mut normal = GeminiProviderState::default();
let mut observer = GeminiProviderState::terminal_observation();
let mut record = observation_record(Vec::new(), Some(reason));
record["candidates"][0]["content"] = content;
record["candidates"][0]["finishMessage"] = json!("provider failure details");
let frames = assert_terminal_record_matches(&mut normal, &mut observer, record);
assert!(frames.iter().any(|frame| matches!(
&frame.event,
CanonicalStreamEvent::UnknownEvent(payload)
if payload["response"]["error"]["code"] == reason
&& payload["response"]["error"]["message"] == "provider failure details"
&& payload["response"]["usage"]["input_tokens"] == 22
)));
assert!(observer.finish(&json!({})).expect("failed EOF").is_empty());
assert_observer_has_no_content_buffers(&observer);
}
}
}
#[test]
fn gemini_provider_state_emits_unknown_events_for_unknown_parts() {
let mut state = GeminiProviderState::default();
@@ -1,6 +1,7 @@
use std::collections::{BTreeMap, BTreeSet};
use serde_json::{json, Map, Value};
use sha2::{Digest, Sha256};
use crate::formats::openai::namespace::NamespaceToolAliases;
use crate::formats::openai::responses::{
@@ -33,6 +34,7 @@ struct OpenAIChatProviderToolState {
#[derive(Default)]
pub struct OpenAIChatProviderState {
terminal_only: bool,
response_id: Option<String>,
model: Option<String>,
actual_service_tier: Option<String>,
@@ -59,6 +61,7 @@ struct OpenAIResponsesProviderToolResultState {
#[derive(Default)]
pub struct OpenAIResponsesProviderState {
terminal_only: bool,
response_id: Option<String>,
model: Option<String>,
actual_service_tier: Option<String>,
@@ -71,11 +74,24 @@ pub struct OpenAIResponsesProviderState {
tool_results: BTreeMap<usize, OpenAIResponsesProviderToolResultState>,
tool_index_by_key: BTreeMap<String, usize>,
image_item_keys: BTreeSet<String>,
opaque_completed_item_keys: BTreeSet<String>,
opaque_completed_item_keys: BTreeSet<OpenAIResponsesOutputItemKey>,
last_tool_index: Option<usize>,
}
#[derive(PartialEq, Eq, PartialOrd, Ord)]
enum OpenAIResponsesOutputItemKey {
Full(String),
Digest([u8; 32]),
}
impl OpenAIChatProviderState {
pub(crate) fn terminal_observation() -> Self {
Self {
terminal_only: true,
..Self::default()
}
}
pub(crate) fn actual_service_tier(&self) -> Option<&str> {
self.actual_service_tier.as_deref()
}
@@ -243,12 +259,14 @@ impl OpenAIChatProviderState {
recognized_delta = true;
if !content.is_empty() {
self.ensure_started(report_context, &mut out);
let (id, model) = self.identity(report_context);
out.push(CanonicalStreamFrame {
id,
model,
event: CanonicalStreamEvent::TextDelta(content.to_string()),
});
if !self.terminal_only {
let (id, model) = self.identity(report_context);
out.push(CanonicalStreamFrame {
id,
model,
event: CanonicalStreamEvent::TextDelta(content.to_string()),
});
}
}
} else if delta.contains_key("content") {
recognized_delta = true;
@@ -258,12 +276,16 @@ impl OpenAIChatProviderState {
recognized_delta = true;
if !reasoning_content.is_empty() {
self.ensure_started(report_context, &mut out);
let (id, model) = self.identity(report_context);
out.push(CanonicalStreamFrame {
id,
model,
event: CanonicalStreamEvent::ReasoningDelta(reasoning_content.to_string()),
});
if !self.terminal_only {
let (id, model) = self.identity(report_context);
out.push(CanonicalStreamFrame {
id,
model,
event: CanonicalStreamEvent::ReasoningDelta(
reasoning_content.to_string(),
),
});
}
}
} else if delta.contains_key("reasoning_content") {
recognized_delta = true;
@@ -272,68 +294,71 @@ impl OpenAIChatProviderState {
if let Some(tool_calls) = delta.get("tool_calls").and_then(Value::as_array) {
recognized_delta = true;
self.ensure_started(report_context, &mut out);
let (id, model) = self.identity(report_context);
for tool_call in tool_calls {
let Some(tool_call_object) = tool_call.as_object() else {
continue;
};
let index = tool_call_object
.get("index")
.and_then(Value::as_u64)
.map(|value| value as usize)
.unwrap_or(0);
let state = self.tool_calls.entry(index).or_default();
if let Some(call_id) = tool_call_object.get("id").and_then(Value::as_str) {
state.id = Some(call_id.to_string());
}
let mut arguments = None;
if let Some(function) =
tool_call_object.get("function").and_then(Value::as_object)
{
if let Some(name) = function.get("name").and_then(Value::as_str) {
state.name = Some(name.to_string());
if !self.terminal_only {
let (id, model) = self.identity(report_context);
for tool_call in tool_calls {
let Some(tool_call_object) = tool_call.as_object() else {
continue;
};
let index = tool_call_object
.get("index")
.and_then(Value::as_u64)
.map(|value| value as usize)
.unwrap_or(0);
let state = self.tool_calls.entry(index).or_default();
if let Some(call_id) = tool_call_object.get("id").and_then(Value::as_str) {
state.id = Some(call_id.to_string());
}
arguments = function
.get("arguments")
.and_then(Value::as_str)
.filter(|arguments| !arguments.is_empty());
}
if !state.started_emitted {
if let Some(arguments) = arguments {
state.pending_arguments.push_str(arguments);
}
if let (Some(call_id), Some(name)) = (state.id.clone(), state.name.clone())
let mut arguments = None;
if let Some(function) =
tool_call_object.get("function").and_then(Value::as_object)
{
out.push(CanonicalStreamFrame {
id: id.clone(),
model: model.clone(),
event: CanonicalStreamEvent::ToolCallStart {
index,
call_id,
name,
},
});
state.started_emitted = true;
if !state.pending_arguments.is_empty() {
if let Some(name) = function.get("name").and_then(Value::as_str) {
state.name = Some(name.to_string());
}
arguments = function
.get("arguments")
.and_then(Value::as_str)
.filter(|arguments| !arguments.is_empty());
}
if !state.started_emitted {
if let Some(arguments) = arguments {
state.pending_arguments.push_str(arguments);
}
if let (Some(call_id), Some(name)) =
(state.id.clone(), state.name.clone())
{
out.push(CanonicalStreamFrame {
id: id.clone(),
model: model.clone(),
event: CanonicalStreamEvent::ToolCallArgumentsDelta {
event: CanonicalStreamEvent::ToolCallStart {
index,
arguments: std::mem::take(&mut state.pending_arguments),
call_id,
name,
},
});
state.started_emitted = true;
if !state.pending_arguments.is_empty() {
out.push(CanonicalStreamFrame {
id: id.clone(),
model: model.clone(),
event: CanonicalStreamEvent::ToolCallArgumentsDelta {
index,
arguments: std::mem::take(&mut state.pending_arguments),
},
});
}
}
} else if let Some(arguments) = arguments {
out.push(CanonicalStreamFrame {
id: id.clone(),
model: model.clone(),
event: CanonicalStreamEvent::ToolCallArgumentsDelta {
index,
arguments: arguments.to_string(),
},
});
}
} else if let Some(arguments) = arguments {
out.push(CanonicalStreamFrame {
id: id.clone(),
model: model.clone(),
event: CanonicalStreamEvent::ToolCallArgumentsDelta {
index,
arguments: arguments.to_string(),
},
});
}
}
} else if delta.contains_key("tool_calls") {
@@ -389,6 +414,13 @@ impl OpenAIChatProviderState {
}
impl OpenAIResponsesProviderState {
pub(crate) fn terminal_observation() -> Self {
Self {
terminal_only: true,
..Self::default()
}
}
pub(crate) fn actual_service_tier(&self) -> Option<&str> {
self.actual_service_tier.as_deref()
}
@@ -494,6 +526,10 @@ impl OpenAIResponsesProviderState {
if text.is_empty() {
return;
}
if self.terminal_only {
self.ensure_started(report_context, out);
return;
}
self.text_parts.entry(key).or_default().push_str(text);
self.ensure_started(report_context, out);
let (id, model) = self.identity(report_context);
@@ -511,6 +547,12 @@ impl OpenAIResponsesProviderState {
key: String,
text: &str,
) {
if self.terminal_only {
if !text.is_empty() {
self.ensure_started(report_context, out);
}
return;
}
let missing = {
let current = self.text_parts.entry(key).or_default();
let missing = if text.starts_with(current.as_str()) {
@@ -543,6 +585,12 @@ impl OpenAIResponsesProviderState {
out: &mut Vec<CanonicalStreamFrame>,
reasoning: &str,
) {
if self.terminal_only {
if !reasoning.is_empty() {
self.ensure_started(report_context, out);
}
return;
}
let missing = if reasoning.starts_with(&self.reasoning) {
reasoning[self.reasoning.len()..].to_string()
} else if self.reasoning == reasoning {
@@ -573,6 +621,10 @@ impl OpenAIResponsesProviderState {
if text.is_empty() {
return;
}
if self.terminal_only {
self.ensure_started(report_context, out);
return;
}
let missing = {
let current = self.reasoning_parts.entry(summary_index).or_default();
let missing = if text.starts_with(current.as_str()) {
@@ -608,6 +660,9 @@ impl OpenAIResponsesProviderState {
out: &mut Vec<CanonicalStreamFrame>,
index: usize,
) {
if self.terminal_only {
return;
}
let (id, model) = self.identity(report_context);
let Some(state) = self.tool_calls.get_mut(&index) else {
return;
@@ -806,12 +861,13 @@ impl OpenAIResponsesProviderState {
if let Some(name) = incoming_chat_name {
state.name = name;
}
let completed_arguments = item
.get("arguments")
.and_then(Value::as_str)
.unwrap_or_default()
.to_string();
Self::merge_tool_call_arguments(state, &completed_arguments);
if !self.terminal_only {
let completed_arguments = item
.get("arguments")
.and_then(Value::as_str)
.unwrap_or_default();
Self::merge_tool_call_arguments(state, completed_arguments);
}
self.emit_ready_function_call(report_context, out, index);
}
@@ -832,10 +888,14 @@ impl OpenAIResponsesProviderState {
.filter(|value| !value.is_empty())
.unwrap_or("custom_tool")
.to_string();
let arguments = tool_arguments_from_maybe_json_string(
item.get("input").or_else(|| item.get("arguments")),
"input",
);
let arguments = if self.terminal_only {
String::new()
} else {
tool_arguments_from_maybe_json_string(
item.get("input").or_else(|| item.get("arguments")),
"input",
)
};
self.emit_generic_tool_call_item(report_context, out, item, output_index, name, arguments);
}
@@ -852,16 +912,20 @@ impl OpenAIResponsesProviderState {
"shell_call" => "shell",
_ => return,
};
let arguments = tool_arguments_from_named_fields(
item,
&[
"action",
"environment",
"status",
"created_by",
"max_output_length",
],
);
let arguments = if self.terminal_only {
String::new()
} else {
tool_arguments_from_named_fields(
item,
&[
"action",
"environment",
"status",
"created_by",
"max_output_length",
],
)
};
self.emit_generic_tool_call_item(
report_context,
out,
@@ -882,7 +946,11 @@ impl OpenAIResponsesProviderState {
if item.get("type").and_then(Value::as_str) != Some("apply_patch_call") {
return;
}
let arguments = tool_arguments_from_named_fields(item, &["operation", "status"]);
let arguments = if self.terminal_only {
String::new()
} else {
tool_arguments_from_named_fields(item, &["operation", "status"])
};
self.emit_generic_tool_call_item(
report_context,
out,
@@ -903,10 +971,14 @@ impl OpenAIResponsesProviderState {
if item.get("type").and_then(Value::as_str) != Some("computer_call") {
return;
}
let arguments = tool_arguments_from_named_fields(
item,
&["action", "actions", "pending_safety_checks", "status"],
);
let arguments = if self.terminal_only {
String::new()
} else {
tool_arguments_from_named_fields(
item,
&["action", "actions", "pending_safety_checks", "status"],
)
};
self.emit_generic_tool_call_item(
report_context,
out,
@@ -941,10 +1013,20 @@ impl OpenAIResponsesProviderState {
.unwrap_or(state.call_id.as_str())
.to_string();
state.name = name;
Self::merge_tool_call_arguments(state, &arguments);
if !self.terminal_only {
Self::merge_tool_call_arguments(state, &arguments);
}
self.emit_ready_tool_call(report_context, out, index);
}
fn tool_result_content(&self, value: Option<&Value>) -> String {
if self.terminal_only {
String::new()
} else {
openai_tool_result_content_from_value(value)
}
}
fn emit_missing_tool_result(
&mut self,
report_context: &Value,
@@ -955,6 +1037,9 @@ impl OpenAIResponsesProviderState {
content: &str,
) {
self.ensure_started(report_context, out);
if self.terminal_only {
return;
}
let state = self.tool_results.entry(index).or_default();
let missing = if !state.emitted {
content.to_string()
@@ -1005,7 +1090,7 @@ impl OpenAIResponsesProviderState {
Some(format!("function_call_output:{tool_use_id}")),
output_index,
);
let content = openai_tool_result_content_from_value(
let content = self.tool_result_content(
item.get("output")
.or_else(|| item.get("content"))
.or_else(|| item.get("delta")),
@@ -1047,7 +1132,7 @@ impl OpenAIResponsesProviderState {
.to_string();
let index =
self.tool_index_for_key(Some(format!("{item_type}:{tool_use_id}")), output_index);
let content = openai_tool_result_content_from_value(
let content = self.tool_result_content(
item.get("output")
.or_else(|| item.get("content"))
.or_else(|| item.get("delta")),
@@ -1110,6 +1195,24 @@ impl OpenAIResponsesProviderState {
if item.get("type").and_then(Value::as_str) != Some("reasoning") {
return;
}
if self.terminal_only {
if item
.get("summary")
.and_then(Value::as_array)
.is_some_and(|summary| {
summary.iter().any(|part| {
part.get("type").and_then(Value::as_str) == Some("summary_text")
&& part
.get("text")
.and_then(Value::as_str)
.is_some_and(|text| !text.is_empty())
})
})
{
self.ensure_started(report_context, out);
}
return;
}
let mut completed_reasoning = String::new();
for raw_summary in item
.get("summary")
@@ -1159,6 +1262,10 @@ impl OpenAIResponsesProviderState {
if !has_image_payload {
return;
}
if self.terminal_only {
self.ensure_started(report_context, out);
return;
}
let index = output_index.unwrap_or(self.image_item_keys.len());
let key = item
.get("id")
@@ -1180,6 +1287,42 @@ impl OpenAIResponsesProviderState {
});
}
fn retained_output_item_key(&self, item: &Map<String, Value>) -> OpenAIResponsesOutputItemKey {
let item_type = item.get("type").and_then(Value::as_str).unwrap_or_default();
if self.terminal_only {
struct DigestWriter(Sha256);
impl std::io::Write for DigestWriter {
fn write(&mut self, bytes: &[u8]) -> std::io::Result<usize> {
self.0.update(bytes);
Ok(bytes.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
// Hash exactly the normal key's bytes, without retaining or first
// serializing an entire encrypted/opaque output item into a string.
let mut writer = DigestWriter(Sha256::new());
writer.0.update(item_type.as_bytes());
if let Some(item_id) = item.get("id").and_then(Value::as_str) {
writer.0.update(b":id:");
writer.0.update(item_id.as_bytes());
} else if let Some(content) = item.get("encrypted_content").and_then(Value::as_str) {
writer.0.update(b":encrypted_content:");
writer.0.update(content.as_bytes());
} else {
writer.0.update(b":");
serde_json::to_writer(&mut writer, item)
.expect("JSON value serialization into a digest cannot fail");
}
return OpenAIResponsesOutputItemKey::Digest(writer.0.finalize().into());
}
OpenAIResponsesOutputItemKey::Full(Self::output_item_key(item))
}
fn output_item_key(item: &Map<String, Value>) -> String {
let item_type = item.get("type").and_then(Value::as_str).unwrap_or_default();
if let Some(item_id) = item.get("id").and_then(Value::as_str) {
@@ -1290,10 +1433,13 @@ impl OpenAIResponsesProviderState {
}
if final_item {
self.opaque_completed_item_keys
.insert(Self::output_item_key(item));
let key = self.retained_output_item_key(item);
self.opaque_completed_item_keys.insert(key);
}
self.ensure_started(report_context, out);
if self.terminal_only {
return;
}
let (id, model) = self.identity(report_context);
out.push(CanonicalStreamFrame {
id,
@@ -1325,7 +1471,7 @@ impl OpenAIResponsesProviderState {
if !self.emit_output_item(report_context, out, item, Some(output_index), true)
&& !self
.opaque_completed_item_keys
.contains(&Self::output_item_key(item))
.contains(&self.retained_output_item_key(item))
{
out.push(self.unknown_frame(report_context, Value::Object(item.clone())));
}
@@ -1504,6 +1650,9 @@ impl OpenAIResponsesProviderState {
.map(|value| value as usize)
.unwrap_or(0);
self.ensure_started(report_context, &mut out);
if self.terminal_only {
return Ok(out);
}
self.reasoning.push_str(piece);
self.reasoning_parts
.entry(summary_index)
@@ -1543,6 +1692,9 @@ impl OpenAIResponsesProviderState {
);
}
self.ensure_started(report_context, &mut out);
if self.terminal_only {
return Ok(out);
}
let (id, model) = self.identity(report_context);
out.push(CanonicalStreamFrame {
id,
@@ -1595,7 +1747,9 @@ impl OpenAIResponsesProviderState {
.unwrap_or("custom_tool")
.to_string();
}
state.arguments.push_str(delta);
if !self.terminal_only {
state.arguments.push_str(delta);
}
self.emit_ready_tool_call(report_context, &mut out, index);
}
"response.custom_tool_call_input.done" => {
@@ -1623,11 +1777,13 @@ impl OpenAIResponsesProviderState {
.unwrap_or("custom_tool")
.to_string();
}
let arguments = tool_arguments_from_maybe_json_string(
Some(&Value::String(input.to_string())),
"input",
);
Self::merge_tool_call_arguments(state, &arguments);
if !self.terminal_only {
let arguments = tool_arguments_from_maybe_json_string(
Some(&Value::String(input.to_string())),
"input",
);
Self::merge_tool_call_arguments(state, &arguments);
}
self.emit_ready_tool_call(report_context, &mut out, index);
}
"response.function_call_arguments.delta" => {
@@ -1654,7 +1810,9 @@ impl OpenAIResponsesProviderState {
if let Some(call_id) = value.get("call_id").and_then(Value::as_str) {
state.call_id = call_id.to_string();
}
state.arguments.push_str(delta);
if !self.terminal_only {
state.arguments.push_str(delta);
}
self.emit_ready_function_call(report_context, &mut out, index);
}
"response.function_call_arguments.done" => {
@@ -1722,7 +1880,9 @@ impl OpenAIResponsesProviderState {
if let Some(name) = incoming_chat_name {
state.name = name;
}
Self::merge_tool_call_arguments(state, arguments);
if !self.terminal_only {
Self::merge_tool_call_arguments(state, arguments);
}
self.emit_ready_function_call(report_context, &mut out, index);
}
"response.function_call_output.delta" | "response.function_call_output.done" => {
@@ -1743,7 +1903,7 @@ impl OpenAIResponsesProviderState {
Some(format!("function_call_output:{tool_use_id}")),
output_index,
);
let content = openai_tool_result_content_from_value(
let content = self.tool_result_content(
value
.get("delta")
.or_else(|| value.get("output"))
@@ -1781,7 +1941,7 @@ impl OpenAIResponsesProviderState {
.map(|value| value as usize);
let index = self
.tool_index_for_key(Some(format!("{item_type}:{tool_use_id}")), output_index);
let content = openai_tool_result_content_from_value(
let content = self.tool_result_content(
value
.get("delta")
.or_else(|| value.get("output"))
@@ -3754,6 +3914,276 @@ mod tests {
format!("data: {}\n", value).into_bytes()
}
fn terminal_frames(frames: Vec<CanonicalStreamFrame>) -> Vec<CanonicalStreamFrame> {
frames
.into_iter()
.filter(|frame| {
matches!(
frame.event,
CanonicalStreamEvent::Start
| CanonicalStreamEvent::UnknownEvent(_)
| CanonicalStreamEvent::Finish { .. }
)
})
.collect()
}
fn assert_no_responses_observation_content(state: &OpenAIResponsesProviderState) {
assert!(state.text_parts.is_empty());
assert_eq!(state.reasoning.capacity(), 0);
assert!(state.reasoning_parts.is_empty());
assert!(state
.tool_calls
.values()
.all(|tool| tool.arguments.capacity() == 0));
assert!(state.tool_results.is_empty());
assert!(state.image_item_keys.is_empty());
assert!(state
.opaque_completed_item_keys
.iter()
.all(|key| matches!(key, OpenAIResponsesOutputItemKey::Digest(_))));
}
#[test]
fn openai_responses_terminal_observation_does_not_retain_long_stream_content() {
let context = json!({"provider_api_format": "openai:responses"});
let mut observed = OpenAIResponsesProviderState::terminal_observation();
let mut full = OpenAIResponsesProviderState::default();
let initial = json!({
"type": "response.output_item.added", "output_index": 0,
"item": {"type": "function_call", "id": "fc_1", "call_id": "call_1", "name": "read", "arguments": ""}
});
assert_eq!(
observed.push_event(&context, &initial).unwrap(),
terminal_frames(full.push_event(&context, &initial).unwrap())
);
let text = "x".repeat(1024);
let events = [
json!({"type": "response.output_text.delta", "delta": text, "output_index": 3}),
json!({"type": "response.reasoning_summary_text.delta", "delta": text, "summary_index": 0}),
json!({"type": "response.function_call_arguments.delta", "delta": text, "output_index": 0}),
json!({"type": "response.custom_tool_call_input.delta", "delta": text, "output_index": 1}),
json!({"type": "response.function_call_output.delta", "delta": text, "output_index": 2, "call_id": "call_1"}),
];
for iteration in 0..2048 {
for event in &events {
let frames = observed.push_event(&context, event).unwrap();
assert!(frames.is_empty());
if iteration < 32 {
assert_eq!(
frames,
terminal_frames(full.push_event(&context, event).unwrap())
);
}
}
assert_no_responses_observation_content(&observed);
}
assert_eq!(observed.tool_calls.len(), 2);
assert_eq!(observed.tool_index_by_key, full.tool_index_by_key);
assert!(full.text_parts.values().any(|text| text.len() == 32 * 1024));
assert_eq!(full.reasoning.len(), 32 * 1024);
assert_eq!(full.reasoning_parts[&0].len(), 32 * 1024);
assert_eq!(full.tool_calls[&0].arguments.len(), 32 * 1024);
assert!(!full.tool_results.is_empty());
assert_eq!(
observed.finish(&context).unwrap(),
full.finish(&context).unwrap()
);
}
#[test]
fn openai_responses_terminal_observation_skips_completed_content_and_images() {
let context = json!({});
let text = "x".repeat(64 * 1024);
let items = [
json!({"type": "message", "content": [{"type": "output_text", "text": text}]}),
json!({"type": "reasoning", "summary": [{"type": "summary_text", "text": text}]}),
json!({"type": "custom_tool_call", "input": text, "name": "custom"}),
json!({"type": "shell_call", "action": {"command": text}}),
json!({"type": "apply_patch_call", "operation": {"patch": text}}),
json!({"type": "computer_call", "action": {"keys": [text]}}),
json!({"type": "function_call_output", "output": text}),
json!({"type": "custom_tool_call_output", "output": {"text": text}}),
json!({"type": "image_generation_call", "result": text, "status": "completed"}),
];
let mut observed = OpenAIResponsesProviderState::terminal_observation();
let mut full = OpenAIResponsesProviderState::default();
for (index, item) in items.iter().enumerate() {
let event =
json!({"type": "response.output_item.done", "output_index": index, "item": item});
assert_eq!(
observed.push_event(&context, &event).unwrap(),
terminal_frames(full.push_event(&context, &event).unwrap())
);
assert_no_responses_observation_content(&observed);
}
let completed = json!({"type": "response.completed", "response": {
"id": "resp_complete", "model": "model", "service_tier": "Priority",
"output": items, "usage": {"input_tokens": 100, "output_tokens": 20, "input_tokens_details": {"cached_tokens": 0}}
}});
assert_eq!(
observed.push_event(&context, &completed).unwrap(),
terminal_frames(full.push_event(&context, &completed).unwrap())
);
assert_eq!(observed.actual_service_tier(), Some("priority"));
assert_no_responses_observation_content(&observed);
}
#[test]
fn openai_responses_terminal_observation_hashes_opaque_keys_without_changing_deduplication() {
let context = json!({});
let text = "x".repeat(64 * 1024);
let items = [
json!({"type": "future_item", "id": "id:with:separators", "encrypted_content": text}),
json!({"type": "compaction", "encrypted_content": text}),
json!({"type": "future_item", "payload": {"text": text, "escaped": "\n\"\\"}}),
];
let mut observed = OpenAIResponsesProviderState::terminal_observation();
let mut full = OpenAIResponsesProviderState::default();
for item in &items {
let object = item.as_object().unwrap();
let normal_key = OpenAIResponsesProviderState::output_item_key(object);
let expected: [u8; 32] = Sha256::digest(normal_key.as_bytes()).into();
let OpenAIResponsesOutputItemKey::Digest(actual) =
observed.retained_output_item_key(object)
else {
panic!("observation must retain only an opaque item digest");
};
assert_eq!(actual, expected);
let event = json!({"type": "response.output_item.done", "item": item});
for _ in 0..2 {
assert_eq!(
observed.push_event(&context, &event).unwrap(),
terminal_frames(full.push_event(&context, &event).unwrap())
);
}
}
assert_eq!(observed.opaque_completed_item_keys.len(), items.len());
assert_no_responses_observation_content(&observed);
let mut final_items = items.to_vec();
final_items.push(json!({"type": "future_item", "id": "new_item"}));
let final_event =
json!({"type": "response.completed", "response": {"output": final_items}});
let frames = observed.push_event(&context, &final_event).unwrap();
assert_eq!(
frames,
terminal_frames(full.push_event(&context, &final_event).unwrap())
);
assert_eq!(
frames
.iter()
.filter(|frame| matches!(frame.event, CanonicalStreamEvent::UnknownEvent(_)))
.count(),
1
);
}
#[test]
fn openai_responses_terminal_observation_preserves_tool_identity_validation() {
let context = json!({"original_request_body": {"tools": [{
"type": "namespace", "name": "reports", "description": "Reporting tools", "tools": [{
"type": "function", "name": "write_report", "parameters": {"type": "object"}
}]
}]}});
let expected_alias = NamespaceToolAliases::from_report_context(&context)
.chat_name("reports", "write_report")
.expect("fixture must contain a valid namespace tool")
.to_owned();
let mut observed = OpenAIResponsesProviderState::terminal_observation();
let mut full = OpenAIResponsesProviderState::default();
let events = [
json!({"type": "response.function_call_arguments.delta", "item_id": "fc_1", "output_index": 0, "delta": "{"}),
json!({"type": "response.output_item.added", "output_index": 0, "item": {
"type": "function_call", "id": "fc_1", "namespace": "reports", "name": "write_report", "arguments": ""
}}),
json!({"type": "response.function_call_arguments.done", "item_id": "fc_1", "call_id": "call_1", "namespace": "reports", "arguments": "{}"}),
json!({"type": "response.function_call_output.delta", "call_id": "call_1", "output_index": 7, "delta": "result"}),
json!({"type": "response.function_call_arguments.done", "item_id": "fc_other", "name": "ordinary", "arguments": "{}"}),
json!({"type": "response.function_call_arguments.done", "item_id": "fc_1", "namespace": "missing", "arguments": "{}"}),
json!({"type": "response.output_item.done", "item": {
"type": "function_call", "id": "invalid", "name": "read", "caller": {"type": "future"}, "arguments": "{}"
}}),
json!({"type": "response.output_item.done", "item": {"content": "missing type"}}),
json!({"type": "response.future.delta", "payload": "unsupported"}),
];
let mut unknowns = 0;
for (step, event) in events.into_iter().enumerate() {
let frames = observed.push_event(&context, &event).unwrap();
assert_eq!(
frames,
terminal_frames(full.push_event(&context, &event).unwrap())
);
unknowns += frames
.iter()
.filter(|frame| matches!(frame.event, CanonicalStreamEvent::UnknownEvent(_)))
.count();
assert_eq!(observed.tool_index_by_key, full.tool_index_by_key);
assert_eq!(observed.last_tool_index, full.last_tool_index);
assert_eq!(observed.tool_calls.len(), full.tool_calls.len());
for (index, tool) in &observed.tool_calls {
assert_eq!(tool.name, full.tool_calls[index].name);
assert_eq!(tool.call_id, full.tool_calls[index].call_id);
}
if step == 1 || step == 2 {
assert_eq!(unknowns, 0);
assert_eq!(observed.tool_calls[&0].name, expected_alias);
assert_eq!(
observed.tool_calls[&0].call_id,
if step == 1 { "" } else { "call_1" }
);
}
assert_no_responses_observation_content(&observed);
}
assert_eq!(unknowns, 4);
assert_eq!(
observed.finish(&context).unwrap(),
full.finish(&context).unwrap()
);
}
#[test]
fn openai_chat_terminal_observation_does_not_buffer_arguments_before_identity() {
let context = json!({});
let mut observed = OpenAIChatProviderState::terminal_observation();
let mut full = OpenAIChatProviderState::default();
let delta = json!({"choices": [{"delta": {
"content": "x".repeat(1024), "reasoning_content": "r".repeat(1024),
"tool_calls": [{"index": 0, "function": {"arguments": "a".repeat(1024)}}]
}}]});
for iteration in 0..2048 {
let frames = observed
.push_line(&context, data_line(delta.clone()))
.unwrap();
if iteration < 32 {
assert_eq!(
frames,
terminal_frames(full.push_line(&context, data_line(delta.clone())).unwrap())
);
}
assert!(observed.tool_calls.is_empty());
}
assert_eq!(full.tool_calls[&0].pending_arguments.len(), 32 * 1024);
for event in [
json!({"choices": [{"delta": {"tool_calls": [{"index": 0, "id": "call_1", "function": {"name": "read", "arguments": "end"}}]}}]}),
json!({"choices": [{"delta": {"future": "unknown"}}]}),
json!({"choices": [{"delta": {}, "finish_reason": "tool_calls"}]}),
json!({"choices": [], "usage": {"prompt_tokens": 100, "completion_tokens": 20, "prompt_tokens_details": {"cached_tokens": 0}}, "service_tier": "Flex"}),
] {
assert_eq!(
observed
.push_line(&context, data_line(event.clone()))
.unwrap(),
terminal_frames(full.push_line(&context, data_line(event)).unwrap())
);
}
assert_eq!(observed.actual_service_tier(), Some("flex"));
assert!(observed.tool_calls.is_empty());
assert_eq!(
observed.finish(&context).unwrap(),
full.finish(&context).unwrap()
);
}
fn response_sequence_numbers(sse: &str) -> Vec<u64> {
let mut sequence_numbers = Vec::new();
for payload in sse.lines().filter_map(|line| line.strip_prefix("data: ")) {
@@ -423,10 +423,26 @@ impl TerminalStreamParser {
{
return Some(Self::OpenAIImage(OpenAiImageStreamTerminalState::default()));
}
ProviderStreamParser::for_api_format(provider_api_format).map(Self::Standard)
let provider = match ProviderStreamParser::for_api_format(provider_api_format)? {
ProviderStreamParser::OpenAIChat(_) => {
ProviderStreamParser::OpenAIChat(OpenAIChatProviderState::terminal_observation())
}
ProviderStreamParser::OpenAIResponses(_) => ProviderStreamParser::OpenAIResponses(
OpenAIResponsesProviderState::terminal_observation(),
),
ProviderStreamParser::Gemini(_) => {
ProviderStreamParser::Gemini(GeminiProviderState::terminal_observation())
}
provider @ ProviderStreamParser::Claude(_) => provider,
};
Some(Self::Standard(provider))
}
}
#[cfg(test)]
#[path = "terminal_observation_tests.rs"]
mod terminal_observation_tests;
enum ProviderStreamParser {
OpenAIChat(OpenAIChatProviderState),
OpenAIResponses(OpenAIResponsesProviderState),
@@ -0,0 +1,282 @@
use serde_json::{json, Value};
use super::{ProviderStreamParser, StreamingStandardTerminalObserver, TerminalStreamParser};
fn full_observer(context: &Value) -> StreamingStandardTerminalObserver {
StreamingStandardTerminalObserver {
provider: Some(TerminalStreamParser::Standard(
ProviderStreamParser::for_api_format(context["provider_api_format"].as_str().unwrap())
.unwrap(),
)),
..Default::default()
}
}
// Compare every input prefix, including EOF, to catch changes in identity and terminal timing.
fn assert_summaries_match(context: &Value, events: &[Value]) {
let structured = context["provider_api_format"]
.as_str()
.unwrap()
.starts_with("openai:responses");
for end in 0..=events.len() {
let mut compact = StreamingStandardTerminalObserver::default();
let mut full = full_observer(context);
let mut via_event = StreamingStandardTerminalObserver::default();
for (index, event) in events[..end].iter().enumerate() {
let line = format!("data: {event}\n").into_bytes();
compact.push_line(context, line.clone()).unwrap();
full.push_line(context, line).unwrap();
assert_eq!(
compact.latest_summary(),
full.latest_summary(),
"prefix {index}: {event}"
);
if structured {
via_event.push_event(context, event).unwrap();
assert_eq!(via_event.latest_summary(), full.latest_summary());
}
}
let expected = full.finish(context).unwrap();
assert_eq!(compact.finish(context).unwrap(), expected, "EOF at {end}");
if structured {
assert_eq!(via_event.finish(context).unwrap(), expected);
}
}
}
fn context(format: &str) -> Value {
json!({"provider_api_format": format, "client_api_format": format, "mapped_model": "test-model"})
}
fn completed(output: Vec<Value>) -> Value {
json!({"type":"response.completed","response":{
"id":"resp-final","model":"final-model","status":"completed","service_tier":" PRIORITY ",
"output":output,"usage":{"input_tokens":11,"output_tokens":7,"total_tokens":18,
"input_tokens_details":{"cached_tokens":0},"output_tokens_details":{"reasoning_tokens":3}}
}})
}
#[test]
fn responses_content_snapshots_and_deltas_preserve_every_summary() {
let events = vec![
json!({"type":"response.output_text.delta","delta":""}),
json!({"type":"response.output_text.delta","delta":"he","output_index":0}),
json!({"type":"response.created","response":{"id":"late-id","model":"late-model"}}),
json!({"type":"response.output_text.delta","delta":{"text":"hello"},"output_index":0}),
json!({"type":"response.output_text.done","text":"hello","output_index":0}),
json!({"type":"response.content_part.added","content_index":1,"part":{"type":"output_text","text":"another"}}),
json!({"type":"response.content_part.done","content_index":1,"part":{"type":"output_text","text":"another part"}}),
json!({"type":"response.refusal.delta","delta":"refuse"}),
json!({"type":"response.refusal.done","refusal":"refused"}),
json!({"type":"response.audio.transcript.delta","delta":"audio"}),
json!({"type":"response.audio.transcript.done","transcript":"audio transcript"}),
json!({"type":"response.reasoning_summary_text.delta","delta":"think","summary_index":0}),
json!({"type":"response.reasoning_text.delta","delta":" again","summary_index":0}),
json!({"type":"response.reasoning_summary_part.added","summary_index":1,"part":{"type":"summary_text","text":"second"}}),
json!({"type":"response.reasoning_summary_part.done","summary_index":1,"part":{"type":"summary_text","text":"second thought"}}),
json!({"type":"response.reasoning_summary_text.done","summary_index":0,"text":"think again"}),
json!({"type":"response.reasoning_text.done","text":""}),
completed(vec![
json!({"type":"message","content":[{"type":"output_text","text":"hello"},{"type":"refusal","refusal":"refused"}]}),
json!({"type":"reasoning","summary":[{"type":"summary_text","text":"think again"},{"type":"summary_text","text":"second thought"}]}),
]),
];
for format in ["openai:responses", "openai:responses:compact"] {
assert_summaries_match(&context(format), &events);
for event in &events {
assert_summaries_match(&context(format), std::slice::from_ref(event));
}
}
}
#[test]
fn responses_tools_and_unknown_execution_fields_preserve_summaries() {
let calls = vec![
json!({"type":"function_call","call_id":"call-0","name":"lookup","arguments":"{\"x\":1}"}),
json!({"type":"custom_tool_call","call_id":"call-1","name":"custom","input":"raw input"}),
json!({"type":"shell_call","call_id":"call-2","action":{"commands":["pwd"]}}),
json!({"type":"local_shell_call","call_id":"call-3","action":{"command":["pwd"]}}),
json!({"type":"apply_patch_call","call_id":"call-4","operation":{"type":"update_file","path":"a","diff":"+b"}}),
json!({"type":"computer_call","call_id":"call-5","action":{"type":"click","x":1,"y":2}}),
json!({"type":"function_call","call_id":"bad-caller","name":"lookup","arguments":"{}","caller":{"type":"direct"}}),
json!({"type":"function_call","call_id":"bad-namespace","namespace":42,"name":"lookup","arguments":"{}"}),
];
let mut events = vec![
json!({"type":"response.function_call_arguments.delta","output_index":0,"delta":"{"}),
json!({"type":"response.function_call_arguments.delta","output_index":0,"call_id":"call-0","delta":"\"x\":1}"}),
json!({"type":"response.function_call_arguments.done","output_index":0,"item":{"call_id":"call-0","name":"lookup","arguments":"{\"x\":1}"}}),
json!({"type":"response.function_call_arguments.done","output_index":0,"namespace":{},"arguments":"{}"}),
json!({"type":"response.custom_tool_call_input.delta","output_index":1,"delta":"raw"}),
json!({"type":"response.custom_tool_call_input.done","output_index":1,"input":"raw input"}),
];
for (index, item) in calls.iter().enumerate() {
events.push(json!({"type":"response.output_item.added","output_index":index,"item":item}));
events.push(json!({"type":"response.output_item.done","output_index":index,"item":item}));
}
for kind in [
"function_call",
"custom_tool_call",
"shell_call",
"local_shell_call",
"apply_patch_call",
"computer_call",
] {
events.push(json!({"type":format!("response.{kind}_output.delta"),"call_id":"result","delta":"result"}));
events.push(json!({"type":format!("response.{kind}_output.done"),"call_id":"result","output":{"ok":true}}));
events.push(json!({"type":"response.output_item.done","item":{"type":format!("{kind}_output"),"call_id":"result","output":"result complete"}}));
}
events.push(completed(calls));
assert_summaries_match(&context("openai:responses"), &events);
}
#[test]
fn responses_opaque_dedup_and_images_preserve_unknown_counts() {
let items = vec![
json!({"type":"future_item","id":"stable-id","payload":"first"}),
json!({"type":"future_item","encrypted_content":"encrypted","payload":"second"}),
json!({"type":"future_item","payload":{"no_id":true}}),
json!({"type":"image_generation_call","result":"image data","status":"completed"}),
json!({"type":"reasoning","summary":[],"encrypted_content":"reasoning"}),
];
let mut events = Vec::new();
for item in &items {
events.push(json!({"type":"response.output_item.added","item":item}));
events.push(json!({"type":"response.output_item.done","item":item}));
events.push(json!({"type":"response.output_item.done","item":item}));
}
events.push(json!({"type":"response.output_item.done","item":{"missing_type":true}}));
let mut output = items;
output[0]["payload"] = json!("changed body, same id");
output[1]["payload"] = json!("changed body, same encrypted content");
output.push(json!({"type":"new_future_item","payload":"never emitted"}));
events.push(completed(output));
let ctx = context("openai:responses");
assert_summaries_match(&ctx, &events);
let mut observer = StreamingStandardTerminalObserver::default();
for event in &events {
observer.push_event(&ctx, event).unwrap();
}
assert_eq!(
observer.finish(&ctx).unwrap().unwrap().unknown_event_count,
2
);
}
#[test]
fn responses_namespace_validation_keeps_existing_tool_identity() {
let mut ctx = context("openai:responses");
ctx["original_request_body"] = json!({"tools":[{
"type":"namespace","name":"search","description":"Search tools","tools":[{"type":"function","name":"lookup","parameters":{"type":"object"}}]
}]});
let events = [
json!({"type":"response.output_item.added","output_index":0,"item":{"type":"function_call","call_id":"call","namespace":"search","name":"lookup","arguments":"{"}}),
json!({"type":"response.function_call_arguments.delta","output_index":0,"delta":"}"}),
json!({"type":"response.function_call_arguments.done","output_index":0,"namespace":"search","arguments":"{}"}),
json!({"type":"response.function_call_arguments.done","output_index":0,"namespace":"missing","arguments":"{}"}),
completed(vec![
json!({"type":"function_call","call_id":"call","namespace":"search","name":"lookup","arguments":"{}"}),
]),
];
assert_summaries_match(&ctx, &events);
let mut observer = StreamingStandardTerminalObserver::default();
for event in &events {
observer.push_event(&ctx, event).unwrap();
}
let summary = observer.finish(&ctx).unwrap().unwrap();
assert_eq!(summary.unknown_event_count, 1);
assert_eq!(summary.finish_reason.as_deref(), Some("tool_calls"));
}
#[test]
fn responses_errors_zero_usage_and_event_type_lines_preserve_summaries() {
for terminal in [
json!({"type":"response.failed","response":{"status":"failed","error":{"message":"failed"},"usage":{"input_tokens":0,"output_tokens":0}}}),
json!({"type":"error","error":{"type":"server_error","message":"failed"}}),
json!({"type":"response.incomplete","response":{"status":"incomplete","incomplete_details":{"reason":"max_output_tokens"},"usage":{"input_tokens":2,"output_tokens":0}}}),
json!({"type":"response.done","response":{"status":"completed","usage":{"input_tokens":0,"output_tokens":0,"total_tokens":0}}}),
json!({"type":"response.completed","response":null}),
] {
let ctx = context("openai:responses");
let events = vec![
json!({"type":"response.future"}),
json!({"type":"ping"}),
terminal,
];
assert_summaries_match(&ctx, &events);
let mut compact = StreamingStandardTerminalObserver::default();
let mut full = full_observer(&ctx);
for mut event in events {
let kind = event.as_object_mut().unwrap().remove("type").unwrap();
for line in [
format!("event: {}\n", kind.as_str().unwrap()),
format!("data: {event}\n"),
"\n".to_string(),
] {
compact.push_line(&ctx, line.as_bytes().to_vec()).unwrap();
full.push_line(&ctx, line.into_bytes()).unwrap();
assert_eq!(compact.latest_summary(), full.latest_summary());
}
}
assert_eq!(compact.finish(&ctx).unwrap(), full.finish(&ctx).unwrap());
}
}
#[test]
fn chat_delayed_tool_identity_and_usage_only_preserve_summaries() {
let ctx = context("openai:chat");
assert_summaries_match(
&ctx,
&[
json!({"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{"}}]}}]}),
json!({"id":"late-id","model":"late-model","choices":[{"delta":{"tool_calls":[{"index":0,"id":"call","function":{"name":"lookup","arguments":"}"}}]}}]}),
json!({"choices":[{"delta":{"content":"text","reasoning_content":"reason"}}],"service_tier":"priority"}),
json!({"choices":[{"delta":{},"finish_reason":"tool_calls"}]}),
json!({"choices":[],"usage":{"prompt_tokens":3,"completion_tokens":2,"total_tokens":5}}),
],
);
for event in [
json!({"usage":{"prompt_tokens":0,"completion_tokens":0,"total_tokens":0}}),
json!({"choices":[{"delta":{"tool_calls":"malformed"}}]}),
json!({"choices":[{"delta":{"tool_calls":[null,{}, {"function":null}]}}]}),
json!({"choices":[{"delta":{},"finish_reason":"future_reason"}]}),
json!({"choices":[{"delta":{"future_content":true}}]}),
] {
assert_summaries_match(&ctx, &[event]);
}
}
#[test]
fn gemini_content_tools_and_errors_preserve_summaries() {
let ctx = context("gemini:generate_content");
let parts = vec![
json!({"text":"text"}),
json!({"text":"reason","thought":true,"thoughtSignature":"sig"}),
json!({"functionCall":{"id":"call","name":"lookup","args":{"x":1}}}),
json!({"functionResponse":{"id":"call","name":"lookup","response":{"ok":true}}}),
json!({"inlineData":{"mimeType":"image/png","data":"aW1hZ2U="}}),
json!({"futureContent":"unknown"}),
];
let mut events = Vec::new();
for part in &parts {
let event = json!({"candidates":[{"content":{"parts":[part]}}]});
events.push(event.clone());
events.push(event.clone());
assert_summaries_match(&ctx, &[event]);
}
events.push(json!({"responseId":"late-id","modelVersion":"late-model","candidates":[{"content":{"parts":parts},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":5,"candidatesTokenCount":3,"totalTokenCount":8}}));
assert_summaries_match(&ctx, &events);
for reason in [
"MALFORMED_FUNCTION_CALL",
"SAFETY",
"MAX_TOKENS",
"FUTURE_REASON",
] {
assert_summaries_match(
&ctx,
&[
json!({"response":{"candidates":[{"content":{"parts":[{"text":"partial"}]}}]}}),
json!({"candidates":[{"content":{"parts":[]},"finishReason":reason}],"usageMetadata":{"promptTokenCount":0,"candidatesTokenCount":0}}),
],
);
}
}
+3
View File
@@ -14,3 +14,6 @@ serde.workspace = true
serde_json.workspace = true
thiserror.workspace = true
tokio.workspace = true
[dev-dependencies]
aether-runtime-state.workspace = true
+539 -1
View File
@@ -500,9 +500,14 @@ fn build_settlement_snapshot(
#[cfg(test)]
mod tests {
use std::sync::Arc;
use aether_data_contracts::repository::billing::StoredBillingModelContext;
use aether_data_contracts::repository::usage::UsageBodyCaptureState;
use aether_usage_runtime::{UsageEvent, UsageEventData, UsageEventType};
use aether_runtime_state::{MemoryRuntimeStateConfig, RuntimeState};
use aether_usage_runtime::{
UsageEvent, UsageEventData, UsageEventType, UsageQueue, UsageRuntimeConfig,
};
use async_trait::async_trait;
use serde_json::json;
use serde_json::Value;
@@ -539,6 +544,539 @@ mod tests {
}
}
fn wire_billing_lookup(pricing: Option<Value>, request_price: Option<f64>) -> TestLookup {
TestLookup {
name_context: Some(
StoredBillingModelContext::new(
"provider-wire".to_string(),
Some("pay_as_you_go".to_string()),
Some("key-wire".to_string()),
None,
Some(5),
"global-model-wire".to_string(),
"wire-model".to_string(),
None,
request_price,
pricing,
Some("model-wire".to_string()),
Some("wire-model".to_string()),
None,
None,
None,
)
.expect("wire billing context"),
),
model_id_context: None,
}
}
fn wire_billing_event(request_id: &str) -> UsageEvent {
UsageEvent::new(
UsageEventType::Completed,
request_id,
UsageEventData {
user_id: Some("user-wire".to_string()),
api_key_id: Some("key-wire".to_string()),
provider_name: "OpenAI".to_string(),
provider_id: Some("provider-wire".to_string()),
provider_api_key_id: Some("key-wire".to_string()),
model: "gpt-5.6-sol".to_string(),
target_model: Some("gpt-5.6-sol".to_string()),
request_type: Some("chat".to_string()),
api_format: Some("openai:responses".to_string()),
endpoint_api_format: Some("openai:responses".to_string()),
input_tokens: Some(1_000),
output_tokens: Some(100),
total_tokens: Some(1_100),
cache_creation_input_tokens: Some(0),
cache_creation_ephemeral_5m_input_tokens: Some(0),
cache_creation_ephemeral_1h_input_tokens: Some(0),
cache_read_input_tokens: Some(0),
status_code: Some(200),
first_byte_time_ms: Some(12),
response_time_ms: Some(30),
request_headers: Some(json!({"x-audit": "request"})),
provider_request_headers: Some(json!({"x-audit": "provider request"})),
response_headers: Some(json!({"x-audit": "provider response"})),
client_response_headers: Some(json!({"x-audit": "client response"})),
provider_request_body: Some(json!({"model": "gpt-5.6-sol"})),
provider_request_body_state: Some(UsageBodyCaptureState::Inline),
response_body: Some(json!({"service_tier": "default", "output": "x".repeat(8192)})),
response_body_state: Some(UsageBodyCaptureState::Inline),
request_metadata: Some(json!({
"usage_available": true,
"usage_pricing_available": true,
"api_key_is_standalone": true,
"plan_usage_reservation_token": "550e8400-e29b-41d4-a716-446655440000",
"plan_usage_reservation_deferred": true
})),
..UsageEventData::default()
},
)
}
fn wire_billing_result(event: &UsageEvent) -> Value {
let mut data = serde_json::to_value(&event.data).expect("serialized billing event");
let object = data.as_object_mut().expect("event data object");
for key in [
"request_body",
"provider_request_body",
"response_body",
"client_response_body",
"request_body_state",
"provider_request_body_state",
"response_body_state",
"client_response_body_state",
"request_headers",
"provider_request_headers",
"response_headers",
"client_response_headers",
] {
object.remove(key);
}
if let Some(metadata) = object
.get_mut("request_metadata")
.and_then(Value::as_object_mut)
{
metadata.retain(|key, _| {
matches!(
key.as_str(),
"usage_available"
| "usage_pricing_available"
| "api_key_is_standalone"
| "plan_usage_reservation_token"
| "plan_usage_reservation_deferred"
| "cancelled_request_fee"
| "dimensions"
| "billing_dimensions"
| "billing_snapshot"
| "settlement_snapshot"
| "rate_multiplier"
| "is_free_tier"
| "settlement_snapshot_schema_version"
)
});
for key in ["billing_snapshot", "settlement_snapshot"] {
if let Some(snapshot) = metadata.get_mut(key).and_then(Value::as_object_mut) {
snapshot.remove("calculated_at");
}
}
}
json!({
"event_type": event.event_type,
"request_id": event.request_id,
"timestamp_ms": event.timestamp_ms,
"data": data,
})
}
async fn assert_wire_billing_equivalent(
lookup: &TestLookup,
original: UsageEvent,
) -> UsageEvent {
const LIMIT: usize = 4096;
let original_fields = original.to_stream_fields().expect("legacy full envelope");
assert!(original_fields["payload"].len() > LIMIT);
let provider_body_present = original
.data
.provider_request_body
.as_ref()
.is_some_and(|body| !body.is_null());
let queue = UsageQueue::new(
Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())),
UsageRuntimeConfig {
enabled: true,
queue_payload_max_bytes: LIMIT,
consumer_block_ms: 1,
..UsageRuntimeConfig::default()
},
)
.expect("bounded billing queue");
queue.ensure_consumer_group().await.expect("billing group");
queue
.enqueue(&original)
.await
.expect("diagnostic projection should fit");
let entries = queue
.read_group("billing-wire-reader")
.await
.expect("billing queue read");
assert_eq!(entries.len(), 1);
assert!(entries[0].fields["payload"].len() <= LIMIT);
// The legacy consumer also prepares request facts after decoding the full envelope.
let mut original =
UsageEvent::from_stream_fields(&original_fields).expect("legacy consumer event");
let mut queued =
UsageEvent::from_stream_fields(&entries[0].fields).expect("projected event");
assert!(original.data.response_body.is_some());
assert_eq!(
original.data.provider_request_body.is_some(),
provider_body_present
);
assert!(queued.data.response_body.is_none());
assert!(queued.data.request_headers.is_none());
assert!(queued.data.provider_request_headers.is_none());
assert!(queued.data.response_headers.is_none());
assert!(queued.data.client_response_headers.is_none());
assert_eq!(
queued.data.response_body_state,
Some(UsageBodyCaptureState::Truncated)
);
enrich_usage_event_with_billing(lookup, &mut original)
.await
.expect("original billing");
enrich_usage_event_with_billing(lookup, &mut queued)
.await
.expect("projected billing");
assert_eq!(wire_billing_result(&queued), wire_billing_result(&original));
queued
}
#[tokio::test]
async fn wire_projection_preserves_openai_requested_tier_and_effective_cache_ttl() {
let lookup = wire_billing_lookup(
Some(json!({
"tiers": [{"up_to": null, "input_price_per_1m": 5.0,
"output_price_per_1m": 30.0, "cache_creation_price_per_1m": 6.25,
"cache_read_price_per_1m": 0.5,
"cache_ttl_pricing": [{"ttl_minutes": 60,
"cache_creation_price_per_1m": 100.0, "cache_read_price_per_1m": 100.0}]}],
"processing_tiers": {"priority": {"price_multiplier": 2.0}}
})),
None,
);
let mut event = wire_billing_event("wire-openai-tier");
event.data.cache_creation_input_tokens = Some(100);
event.data.provider_request_body = Some(json!({
"model": "gpt-5.6-sol", "service_tier": "priority", "reasoning": {"effort": "high"}
}));
event.data.request_metadata.as_mut().unwrap()["provider_service_tier"] = json!("flex");
event.data.request_metadata.as_mut().unwrap()["provider_actual_service_tier"] =
json!("flex");
let queued = assert_wire_billing_equivalent(&lookup, event).await;
let metadata = queued.data.request_metadata.as_ref().unwrap();
assert_eq!(metadata["provider_service_tier"], "priority");
assert_eq!(metadata["provider_actual_service_tier"], "flex");
assert_eq!(metadata["provider_reasoning_effort"], "high");
assert_eq!(metadata["billing_dimensions"]["cache_ttl_minutes"], 30);
assert_eq!(
metadata["billing_dimensions"]["billing_processing_tier"],
"priority"
);
assert!(queued.data.total_cost_usd.unwrap() > 0.0);
}
#[tokio::test]
async fn wire_projection_preserves_non_object_body_authority_and_null_decode_semantics() {
let lookup = wire_billing_lookup(
Some(json!({
"tiers": [{"up_to": null, "input_price_per_1m": 5.0,
"output_price_per_1m": 30.0, "cache_creation_price_per_1m": 6.25,
"cache_read_price_per_1m": 0.5,
"cache_ttl_pricing": [{"ttl_minutes": 60,
"cache_creation_price_per_1m": 100.0, "cache_read_price_per_1m": 100.0}]}],
"processing_tiers": {"priority": {"price_multiplier": 2.0}}
})),
None,
);
for (kind, body) in [
("string", json!("not an object")),
("array", json!([{"service_tier": "flex"}])),
("number", json!(42)),
("boolean", json!(false)),
("null", Value::Null),
] {
for state in [
Some(UsageBodyCaptureState::Inline),
Some(UsageBodyCaptureState::Reference),
None,
] {
let mut event = wire_billing_event(&format!("wire-{kind}-{state:?}"));
event.data.input_tokens = Some(1_000_000);
event.data.output_tokens = Some(0);
event.data.total_tokens = Some(1_000_000);
event.data.cache_creation_input_tokens = Some(1_000_000);
event.data.provider_request_body = Some(body.clone());
event.data.provider_request_body_state = state;
let metadata = event.data.request_metadata.as_mut().unwrap();
metadata["provider_service_tier"] = json!("priority");
metadata["provider_reasoning_effort"] = json!("high");
metadata["provider_cache_ttl_minutes"] = json!(60);
let queued = assert_wire_billing_equivalent(&lookup, event).await;
let metadata = queued.data.request_metadata.as_ref().unwrap();
assert!(queued.data.provider_request_body.is_none());
assert_eq!(metadata["provider_cache_ttl_minutes"], 60);
assert_eq!(metadata["billing_dimensions"]["cache_ttl_minutes"], 60);
let expected_tier = if body.is_null() && state.is_some() {
Some("priority")
} else {
None
};
assert_eq!(
usage_event_processing_tiers(&queued.data)
.requested
.as_deref(),
expected_tier,
"{kind} with {state:?}"
);
assert_eq!(
queued.data.total_cost_usd,
Some(if expected_tier.is_some() {
200.0
} else {
100.0
}),
"{kind} with {state:?}"
);
if body.is_null() {
// Option<Value> decodes JSON null as absent, so the old capture marker remains.
assert_eq!(queued.data.provider_request_body_state, state);
assert_eq!(metadata["provider_service_tier"], "priority");
assert_eq!(metadata["provider_reasoning_effort"], "high");
} else {
assert_eq!(
queued.data.provider_request_body_state,
Some(UsageBodyCaptureState::Truncated)
);
assert!(metadata.get("provider_service_tier").is_none());
assert!(metadata.get("provider_reasoning_effort").is_none());
}
}
}
}
#[tokio::test]
async fn wire_projection_preserves_raw_body_ttl_with_non_authoritative_capture_states() {
let lookup = wire_billing_lookup(
Some(json!({
"tiers": [{"up_to": null, "input_price_per_1m": 5.0,
"output_price_per_1m": 30.0, "cache_creation_price_per_1m": 6.25,
"cache_read_price_per_1m": 0.5,
"cache_ttl_pricing": [{"ttl_minutes": 60,
"cache_creation_price_per_1m": 100.0, "cache_read_price_per_1m": 100.0}]}]
})),
None,
);
for state in [
UsageBodyCaptureState::Disabled,
UsageBodyCaptureState::Unavailable,
UsageBodyCaptureState::Truncated,
] {
let mut event = wire_billing_event(&format!("wire-capture-state-{state:?}"));
event.data.input_tokens = Some(1_000_000);
event.data.output_tokens = Some(0);
event.data.total_tokens = Some(1_000_000);
event.data.cache_creation_input_tokens = Some(1_000_000);
event.data.provider_request_body = Some(json!({
"model": "gpt-5.6-sol", "prompt_cache_options": {"ttl": "30m"}
}));
event.data.provider_request_body_state = Some(state);
event.data.request_metadata.as_mut().unwrap()["provider_cache_ttl_minutes"] = json!(60);
let queued = assert_wire_billing_equivalent(&lookup, event).await;
assert_eq!(queued.data.provider_request_body_state, Some(state));
assert!(queued.data.provider_request_body.is_none());
assert_eq!(queued.data.total_cost_usd, Some(6.25));
assert_eq!(
queued.data.request_metadata.as_ref().unwrap()["billing_dimensions"]
["cache_ttl_minutes"],
30
);
}
}
#[tokio::test]
async fn wire_projection_rejects_typed_none_when_omitting_raw_ttl_would_change_billing() {
let lookup = wire_billing_lookup(
Some(json!({
"tiers": [{"up_to": null, "input_price_per_1m": 5.0,
"output_price_per_1m": 30.0, "cache_creation_price_per_1m": 6.25,
"cache_read_price_per_1m": 0.5,
"cache_ttl_pricing": [{"ttl_minutes": 60,
"cache_creation_price_per_1m": 100.0, "cache_read_price_per_1m": 100.0}]}]
})),
None,
);
let mut event = wire_billing_event("wire-typed-none-ttl");
event.data.input_tokens = Some(1_000_000);
event.data.output_tokens = Some(0);
event.data.total_tokens = Some(1_000_000);
event.data.cache_creation_input_tokens = Some(1_000_000);
event.data.provider_request_body = Some(json!({
"model": "gpt-5.6-sol", "prompt_cache_options": {"ttl": "30m"}
}));
event.data.provider_request_body_state = Some(UsageBodyCaptureState::None);
event.data.request_metadata.as_mut().unwrap()["provider_cache_ttl_minutes"] = json!(60);
let original_fields = event.to_stream_fields().expect("legacy full envelope");
assert!(original_fields["payload"].len() > 4096);
let mut legacy = UsageEvent::from_stream_fields(&original_fields).expect("legacy consumer");
assert_eq!(
legacy.data.provider_request_body_state,
Some(UsageBodyCaptureState::None)
);
assert!(legacy.data.provider_request_body.is_some());
assert!(legacy
.data
.request_metadata
.as_ref()
.unwrap()
.get("provider_cache_ttl_minutes")
.is_none());
enrich_usage_event_with_billing(&lookup, &mut legacy)
.await
.expect("legacy billing");
assert_eq!(legacy.data.total_cost_usd, Some(6.25));
let queue = UsageQueue::new(
Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())),
UsageRuntimeConfig {
enabled: true,
queue_payload_max_bytes: 4096,
consumer_block_ms: 1,
..UsageRuntimeConfig::default()
},
)
.expect("bounded billing queue");
queue.ensure_consumer_group().await.expect("billing group");
assert!(matches!(
queue.enqueue(&event).await,
Err(aether_data_contracts::DataLayerError::InvalidInput(_))
));
assert_eq!(event.to_stream_fields().unwrap(), original_fields);
assert!(queue
.read_group("typed-none-reader")
.await
.unwrap()
.is_empty());
// The unchanged source event remains usable by the terminal direct-write fallback.
enrich_usage_event_with_billing(&lookup, &mut event)
.await
.expect("direct fallback billing");
assert_eq!(wire_billing_result(&event), wire_billing_result(&legacy));
}
#[tokio::test]
async fn wire_projection_preserves_claude_cache_ttl_and_explicit_zero_segments() {
let lookup = wire_billing_lookup(
Some(json!({
"tiers": [{"up_to": null, "input_price_per_1m": 3.0,
"output_price_per_1m": 15.0, "cache_creation_price_per_1m": 3.75,
"cache_read_price_per_1m": 0.3,
"cache_ttl_pricing": [{"ttl_minutes": 60,
"cache_creation_price_per_1m": 6.0, "cache_read_price_per_1m": 0.6}]}]
})),
None,
);
let mut event = wire_billing_event("wire-claude-cache");
event.data.model = "claude-sonnet-4-6".to_string();
event.data.target_model = Some("claude-sonnet-4-6".to_string());
event.data.api_format = Some("claude:chat".to_string());
event.data.endpoint_api_format = Some("claude:chat".to_string());
event.data.provider_request_body = None;
event.data.provider_request_body_state = Some(UsageBodyCaptureState::Disabled);
event.data.cache_creation_input_tokens = Some(200);
event.data.cache_creation_ephemeral_5m_input_tokens = Some(0);
event.data.cache_creation_ephemeral_1h_input_tokens = Some(100);
event.data.request_metadata.as_mut().unwrap()["provider_cache_ttl_minutes"] = json!(60);
let queued = assert_wire_billing_equivalent(&lookup, event).await;
assert_eq!(
queued.data.cache_creation_ephemeral_5m_input_tokens,
Some(0)
);
assert_eq!(queued.data.cache_read_input_tokens, Some(0));
let dimensions = &queued.data.request_metadata.as_ref().unwrap()["billing_dimensions"];
assert_eq!(dimensions["cache_ttl_minutes"], 60);
assert_eq!(dimensions["cache_creation_ephemeral_1h_tokens"], 100);
assert_eq!(dimensions["cache_creation_uncategorized_tokens"], 100);
}
#[tokio::test]
async fn wire_projection_preserves_unknown_zero_error_and_cancellation_billing() {
let lookup = wire_billing_lookup(None, Some(0.02));
for mode in ["unknown", "unpriced", "zero", "error_present", "cancelled"] {
let mut event = wire_billing_event(mode);
event.data.input_tokens = Some(0);
event.data.output_tokens = Some(0);
event.data.total_tokens = Some(0);
match mode {
"unknown" => {
event.data.input_tokens = None;
event.data.output_tokens = None;
event.data.total_tokens = None;
event.data.request_metadata.as_mut().unwrap()["usage_available"] = json!(false);
}
"unpriced" => {
event.data.input_tokens = Some(12);
event.data.output_tokens = Some(3);
event.data.total_tokens = Some(15);
event.data.request_metadata.as_mut().unwrap()["usage_pricing_available"] =
json!(false);
}
"error_present" => event.data.error_message = Some(String::new()),
"cancelled" => event.event_type = UsageEventType::Cancelled,
_ => {}
}
let queued = assert_wire_billing_equivalent(&lookup, event).await;
match mode {
"unknown" => {
assert_eq!(queued.data.input_tokens, None);
assert_eq!(queued.data.total_cost_usd, None);
}
"unpriced" => {
assert_eq!(queued.data.input_tokens, Some(12));
assert_eq!(queued.data.total_cost_usd, None);
}
"error_present" => {
assert_eq!(queued.data.error_message.as_deref(), Some(""));
assert_eq!(queued.data.total_cost_usd, Some(0.0));
}
"zero" | "cancelled" => assert_eq!(queued.data.total_cost_usd, Some(0.02)),
_ => unreachable!(),
}
}
}
#[tokio::test]
async fn wire_projection_preserves_image_matrix_dimensions_and_request_count() {
let lookup = wire_billing_lookup(
Some(json!({
"image_output_price_default": 0.01,
"image_output_prices": {"1536x1024": {"medium": 0.041, "high": 0.165}}
})),
Some(0.02),
);
let mut event = wire_billing_event("wire-image");
event.data.request_type = Some("image".to_string());
event.data.api_format = Some("openai:image".to_string());
event.data.endpoint_api_format = Some("openai:image".to_string());
event.data.input_tokens = Some(0);
event.data.output_tokens = Some(0);
event.data.total_tokens = Some(0);
event.data.request_metadata.as_mut().unwrap()["dimensions"] = json!({
"image_count": 2, "image_size": "1536x1024", "image_quality": "medium",
"image_output_format": "png"
});
let queued = assert_wire_billing_equivalent(&lookup, event).await;
let metadata = queued.data.request_metadata.as_ref().unwrap();
assert_eq!(metadata["billing_dimensions"]["image_count"], 2);
assert_eq!(metadata["billing_dimensions"]["request_count"], 2);
assert_eq!(
metadata["billing_dimensions"]["image_price_key"],
"1536x1024:medium"
);
assert_eq!(
metadata["billing_snapshot"]["cost_breakdown"]["image_output_cost"],
0.082
);
assert_eq!(
metadata["billing_snapshot"]["cost_breakdown"]["request_cost"],
0.04
);
}
#[tokio::test]
async fn unmetered_session_audit_does_not_fabricate_tokens_or_request_cost() {
let lookup = TestLookup {
@@ -67,6 +67,21 @@ GROUP BY
FLOOR(EXTRACT(EPOCH FROM (created_at - TO_TIMESTAMP($2))) / $4)::BIGINT
"#;
const RUNTIME_CANDIDATE_COLUMNS: &str = r#"
SELECT
id, request_id, user_id, api_key_id,
NULL::text AS username, NULL::text AS api_key_name,
candidate_index, retry_index, provider_id, endpoint_id, key_id, status,
NULL::text AS skip_reason, is_cached, status_code,
NULL::text AS error_type, NULL::text AS error_message,
latency_ms, concurrent_requests,
NULL::jsonb AS extra_data, NULL::jsonb AS required_capabilities,
CAST(EXTRACT(EPOCH FROM created_at) * 1000 AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM started_at) * 1000 AS BIGINT) AS started_at_unix_ms,
CAST(EXTRACT(EPOCH FROM finished_at) * 1000 AS BIGINT) AS finished_at_unix_ms
FROM request_candidates
"#;
const UPSERT_SQL_TEMPLATE: &str = r#"
INSERT INTO request_candidates (
id,
@@ -561,12 +576,29 @@ impl SqlxRequestCandidateReadRepository {
pub async fn list_recent(
&self,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
self.list_recent_with_columns(limit, candidate_columns())
.await
}
pub async fn list_recent_runtime(
&self,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
self.list_recent_with_columns(limit, RUNTIME_CANDIDATE_COLUMNS)
.await
}
async fn list_recent_with_columns(
&self,
limit: usize,
columns: &'static str,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Postgres>::new(candidate_columns());
let mut builder = QueryBuilder::<Postgres>::new(columns);
builder.push(" ORDER BY created_at DESC");
push_limit(
&mut builder,
@@ -1068,6 +1100,13 @@ impl RequestCandidateReadRepository for SqlxRequestCandidateReadRepository {
Self::list_recent(self, limit).await
}
async fn list_recent_runtime(
&self,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
Self::list_recent_runtime(self, limit).await
}
async fn list_finalized_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
@@ -1699,4 +1738,67 @@ VALUES ($1, $2, 0, 0, 'pending', $3, $4, $5::json, $6::json, $7, NOW())
.await
.expect("candidate NUL test rows should clean up");
}
#[tokio::test]
#[ignore = "requires isolated AETHER_TEST_DATABASE_URL; uses a connection-local table"]
async fn live_postgres_candidate_runtime_projection_preserves_metadata_and_admin_rows() {
let database_url =
std::env::var("AETHER_TEST_DATABASE_URL").expect("isolated test database");
let pool = sqlx::postgres::PgPoolOptions::new()
.max_connections(1)
.connect(&database_url)
.await
.unwrap();
sqlx::query(
r#"
CREATE TEMP TABLE request_candidates (
id text, request_id text, user_id text, api_key_id text,
username text, api_key_name text, candidate_index integer, retry_index integer,
provider_id text, endpoint_id text, key_id text, status text, skip_reason text,
is_cached boolean, status_code integer, error_type text, error_message text,
latency_ms integer, concurrent_requests integer, extra_data jsonb, required_capabilities jsonb,
created_at timestamptz, started_at timestamptz, finished_at timestamptz
)
"#,
)
.execute(&pool)
.await
.unwrap();
for index in 0..3 {
sqlx::query(
r#"
INSERT INTO request_candidates VALUES (
$1, $2, 'user', 'api-key', NULL, NULL, $3, 0,
'provider', 'endpoint', 'key', 'failed', NULL, false, 500,
'upstream_error', 'admin diagnostic', 20, 17, $4, '{"vision": true}'::jsonb,
TO_TIMESTAMP(100 + $3), TO_TIMESTAMP(101 + $3), TO_TIMESTAMP(102 + $3)
)
"#,
)
.bind(format!("candidate-{index}"))
.bind(format!("request-{index}"))
.bind(index)
.bind(json!({"upstream_response": {"body": "x".repeat(32_768)}}))
.execute(&pool)
.await
.unwrap();
}
let repository = SqlxRequestCandidateReadRepository::new(pool.clone());
let full = repository.list_recent(2).await.unwrap();
let runtime = repository.list_recent_runtime(2).await.unwrap();
assert_eq!(
runtime,
full.iter()
.map(|row| row.runtime_snapshot())
.collect::<Vec<_>>()
);
assert_eq!(runtime[0].id, "candidate-2");
assert_eq!(runtime[0].concurrent_requests, Some(17));
assert!(runtime[0].extra_data.is_none());
assert!(runtime[0].error_message.is_none());
assert!(full[0].extra_data.is_some());
assert_eq!(repository.list_recent(2).await.unwrap(), full);
assert!(repository.list_recent_runtime(0).await.unwrap().is_empty());
pool.close().await;
}
}
@@ -52,7 +52,7 @@ pub use migrations::{
run_migrations_with_bootstrap, BootstrapFuture, PostgresMigrationBootstrap, POSTGRES_MIGRATOR,
};
pub use oauth_providers::SqlxOAuthProviderRepository;
pub use pool::{PostgresPool, PostgresPoolFactory};
pub use pool::{acquire_postgres_migration_connection, PostgresPool, PostgresPoolFactory};
pub use pool_scores::PostgresPoolMemberScoreRepository;
pub use provider_catalog::SqlxProviderCatalogReadRepository;
pub use proxy_nodes::SqlxProxyNodeRepository;
@@ -93,7 +93,7 @@ pub async fn run_migrations_with_bootstrap(
pool: &PgPool,
bootstrap: &dyn PostgresMigrationBootstrap,
) -> Result<(), MigrateError> {
let mut conn = pool.acquire().await?;
let mut conn = crate::pool::acquire_postgres_migration_connection(pool).await?;
if POSTGRES_MIGRATOR.locking {
conn.lock().await?;
@@ -132,7 +132,7 @@ pub async fn prepare_database_for_startup_with_bootstrap(
pool: &PgPool,
bootstrap: &dyn PostgresMigrationBootstrap,
) -> Result<Vec<PendingMigrationInfo>, MigrateError> {
let mut conn = pool.acquire().await?;
let mut conn = crate::pool::acquire_postgres_migration_connection(pool).await?;
if POSTGRES_MIGRATOR.locking {
conn.lock().await?;
@@ -4,6 +4,70 @@ use sqlx::PgPool;
use std::str::FromStr;
use std::time::Duration;
const STATEMENT_TIMEOUT_ENV: &str = "AETHER_GATEWAY_DATA_POSTGRES_STATEMENT_TIMEOUT_MS";
const LOCK_TIMEOUT_ENV: &str = "AETHER_GATEWAY_DATA_POSTGRES_LOCK_TIMEOUT_MS";
#[derive(Debug, Clone, Copy)]
struct PostgresSessionTimeouts {
statement_ms: u32,
lock_ms: u32,
}
impl PostgresSessionTimeouts {
fn from_env() -> Result<Self, DataLayerError> {
Ok(Self {
statement_ms: read_timeout_env(STATEMENT_TIMEOUT_ENV, 30_000)?,
lock_ms: read_timeout_env(LOCK_TIMEOUT_ENV, 3_000)?,
})
}
fn apply(self, options: PgConnectOptions) -> PgConnectOptions {
options.options([
("statement_timeout", self.statement_ms),
("lock_timeout", self.lock_ms),
])
}
}
fn parse_timeout_ms(name: &str, value: &str) -> Result<u32, DataLayerError> {
value
.trim()
.parse::<u32>()
.ok()
.filter(|value| *value <= i32::MAX as u32)
.ok_or_else(|| {
DataLayerError::InvalidConfiguration(format!(
"{name} must be milliseconds in 0..=2147483647 (0 disables the timeout)"
))
})
}
fn read_timeout_env(name: &str, default: u32) -> Result<u32, DataLayerError> {
match std::env::var(name) {
Ok(value) => parse_timeout_ms(name, &value),
Err(std::env::VarError::NotPresent) => Ok(default),
Err(std::env::VarError::NotUnicode(_)) => Err(DataLayerError::InvalidConfiguration(
format!("{name} must contain a valid integer"),
)),
}
}
/// Migration and historical backfill connections are discarded on every exit path,
/// including cancellation, so their relaxed deadlines cannot escape into request work.
pub async fn acquire_postgres_migration_connection(
pool: &PgPool,
) -> Result<sqlx::pool::PoolConnection<sqlx::Postgres>, sqlx::Error> {
let mut conn = pool.acquire().await?;
conn.close_on_drop();
sqlx::query("SET statement_timeout = 0")
.execute(&mut *conn)
.await?;
sqlx::query("SET lock_timeout = 0")
.execute(&mut *conn)
.await?;
Ok(conn)
}
fn connect_options(config: &PostgresPoolConfig) -> Result<PgConnectOptions, DataLayerError> {
config.validate()?;
let options = PgConnectOptions::from_str(config.database_url.trim()).map_err(|err| {
@@ -33,12 +97,16 @@ pub type PostgresPool = PgPool;
#[derive(Debug, Clone)]
pub struct PostgresPoolFactory {
config: PostgresPoolConfig,
timeouts: PostgresSessionTimeouts,
}
impl PostgresPoolFactory {
pub fn new(config: PostgresPoolConfig) -> Result<Self, DataLayerError> {
config.validate()?;
Ok(Self { config })
Ok(Self {
config,
timeouts: PostgresSessionTimeouts::from_env()?,
})
}
pub fn config(&self) -> &PostgresPoolConfig {
@@ -46,7 +114,7 @@ impl PostgresPoolFactory {
}
pub fn connect_lazy(&self) -> Result<PostgresPool, DataLayerError> {
let options = connect_options(&self.config)?;
let options = self.timeouts.apply(connect_options(&self.config)?);
Ok(PgPoolOptions::new()
.min_connections(self.config.min_connections)
.max_connections(self.config.max_connections)
@@ -59,10 +127,171 @@ impl PostgresPoolFactory {
#[cfg(test)]
mod tests {
use super::{connect_options, PostgresPoolFactory};
use super::{connect_options, parse_timeout_ms, PostgresPoolFactory, PostgresSessionTimeouts};
use crate::PostgresPoolConfig;
use sqlx::postgres::PgSslMode;
#[tokio::test]
async fn migration_connection_future_is_send() {
fn assert_send(_: impl Send) {}
let pool = sqlx::postgres::PgPoolOptions::new()
.connect_lazy("postgres://localhost/aether")
.unwrap();
assert_send(super::acquire_postgres_migration_connection(&pool));
}
#[test]
fn validates_session_timeout_milliseconds() {
assert_eq!(parse_timeout_ms("timeout", "0").unwrap(), 0);
assert_eq!(parse_timeout_ms("timeout", " 3000 ").unwrap(), 3_000);
assert_eq!(
parse_timeout_ms("timeout", "2147483647").unwrap(),
i32::MAX as u32
);
for invalid in ["", "-1", "3s", "2147483648", "4294967296"] {
assert!(parse_timeout_ms("timeout", invalid).is_err());
}
}
#[test]
fn session_deadlines_preserve_unrelated_connection_options() {
let options = PostgresSessionTimeouts {
statement_ms: 30_000,
lock_ms: 3_000,
}
.apply(sqlx::postgres::PgConnectOptions::new().options([("search_path", "audit")]));
assert_eq!(
options.get_options(),
Some("-c search_path=audit -c statement_timeout=30000 -c lock_timeout=3000")
);
}
#[tokio::test]
#[ignore = "requires an isolated AETHER_TEST_DATABASE_URL"]
async fn live_session_deadlines_rollback_transactions_and_isolate_migration_overrides() {
use crate::error::SqlxResultExt;
use crate::{PostgresTransactionOptions, PostgresTransactionRunner};
let factory = PostgresPoolFactory {
config: PostgresPoolConfig {
database_url: std::env::var("AETHER_TEST_DATABASE_URL").expect("test database URL"),
min_connections: 0,
max_connections: 2,
..PostgresPoolConfig::default()
},
timeouts: PostgresSessionTimeouts {
statement_ms: 100,
lock_ms: 40,
},
};
let pool = factory.connect_lazy().unwrap();
let table = format!("deadline_test_{}", uuid::Uuid::new_v4().simple());
sqlx::query(&format!(
"CREATE TABLE {table} (id INTEGER PRIMARY KEY, value INTEGER NOT NULL)"
))
.execute(&pool)
.await
.unwrap();
sqlx::query(&format!("INSERT INTO {table} VALUES (1, 0)"))
.execute(&pool)
.await
.unwrap();
let mut blocker = pool.begin().await.unwrap();
sqlx::query(&format!("UPDATE {table} SET value = 7 WHERE id = 1"))
.execute(&mut *blocker)
.await
.unwrap();
let runner = PostgresTransactionRunner::new(pool.clone());
let insert = format!("INSERT INTO {table} VALUES (2, 2)");
let update = format!("UPDATE {table} SET value = 9 WHERE id = 1");
let started = std::time::Instant::now();
let error = runner
.run_read_write(|tx| {
Box::pin(async move {
sqlx::query(&insert)
.execute(&mut **tx)
.await
.map_postgres_err()?;
sqlx::query(&update)
.execute(&mut **tx)
.await
.map_postgres_err()?;
Ok(())
})
})
.await
.unwrap_err();
assert!(error.to_string().contains("SQLSTATE 55P03"), "{error}");
assert!(started.elapsed() < std::time::Duration::from_secs(2));
blocker.rollback().await.unwrap();
assert_eq!(
sqlx::query_scalar::<_, i64>(&format!("SELECT COUNT(*) FROM {table}"))
.fetch_one(&pool)
.await
.unwrap(),
1
);
assert_eq!(
sqlx::query_scalar::<_, i32>(&format!("SELECT value FROM {table} WHERE id = 1"))
.fetch_one(&pool)
.await
.unwrap(),
0
);
let error = sqlx::query("SELECT pg_sleep(0.3)")
.execute(&pool)
.await
.map_postgres_err()
.unwrap_err();
assert!(error.to_string().contains("SQLSTATE 57014"), "{error}");
runner
.run(
PostgresTransactionOptions {
statement_timeout_ms: Some(1_000),
..PostgresTransactionOptions::read_write()
},
|tx| {
Box::pin(async move {
sqlx::query("SELECT pg_sleep(0.15)")
.execute(&mut **tx)
.await
.map_postgres_err()?;
Ok(())
})
},
)
.await
.unwrap();
let mut migration = super::acquire_postgres_migration_connection(&pool)
.await
.unwrap();
sqlx::query("SELECT pg_sleep(0.15)")
.execute(&mut *migration)
.await
.unwrap();
drop(migration);
for _ in 0..2 {
let configured: i64 = sqlx::query_scalar(
"SELECT setting::BIGINT FROM pg_settings WHERE name = 'statement_timeout'",
)
.fetch_one(&pool)
.await
.unwrap();
assert_eq!(
configured, 100,
"relaxed/local overrides must not leak into pooled requests"
);
}
sqlx::query(&format!("DROP TABLE {table}"))
.execute(&pool)
.await
.unwrap();
pool.close().await;
}
fn ssl_mode(url: &str, require_ssl: bool) -> PgSslMode {
connect_options(&PostgresPoolConfig {
database_url: url.to_string(),
@@ -1,5 +1,5 @@
use async_trait::async_trait;
use sqlx::{PgPool, Row};
use sqlx::{PgPool, Postgres, QueryBuilder, Row};
use aether_data_contracts::repository::settlement::{
finite_wallet_available_usd, plan_finite_wallet_debit, settlement_billable_cost_usd,
@@ -297,6 +297,76 @@ fn usage_policy_subject_missing() -> DataLayerError {
DataLayerError::InvalidInput("usage policy subject does not exist".to_string())
}
fn usage_policy_window_aggregate_query(
windows: impl Iterator<Item = (u64, u64)>,
aggregate: &str,
) -> Result<(QueryBuilder<'static, Postgres>, i64, i64), DataLayerError> {
let mut builder = QueryBuilder::new("SELECT ");
let mut earliest = i64::MAX;
let mut latest = i64::MIN;
for (index, (start, end)) in windows.enumerate() {
let start = usage_policy_cost_i64(start, "usage policy window start")?;
let end = usage_policy_cost_i64(end, "usage policy window end")?;
earliest = earliest.min(start);
latest = latest.max(end);
if index > 0 {
builder.push(", ");
}
builder
.push("COALESCE(")
.push(aggregate)
.push(" FILTER (WHERE admitted_at >= TO_TIMESTAMP(")
.push_bind(start)
.push("::double precision) AND admitted_at < TO_TIMESTAMP(")
.push_bind(end)
.push("::double precision)), 0)::BIGINT");
}
Ok((builder, earliest, latest))
}
async fn usage_policy_request_window_counts(
tx: &mut sqlx::Transaction<'_, Postgres>,
input: &ReserveUsagePolicyRequestInput,
) -> Result<sqlx::postgres::PgRow, DataLayerError> {
let (mut query, earliest, latest) = usage_policy_window_aggregate_query(
input
.windows
.iter()
.map(|window| (window.starts_at_unix_secs, window.ends_at_unix_secs)),
"COUNT(*)",
)?;
// The subject lock protects all windows. One bounded history scan replaces
// repeated scans of overlapping windows without approximating their counts.
query
.push(" FROM usage_request_admissions WHERE subject_id = ")
.push_bind(input.subject_id.clone())
.push(" AND state = 'active' AND admitted_at >= TO_TIMESTAMP(")
.push_bind(earliest)
.push("::double precision) AND admitted_at < TO_TIMESTAMP(")
.push_bind(latest)
.push("::double precision)");
query.build().fetch_one(&mut **tx).await.map_postgres_err()
}
async fn usage_policy_cost_window_totals(
tx: &mut sqlx::Transaction<'_, Postgres>,
input: &ReserveUsagePolicyCostInput,
) -> Result<sqlx::postgres::PgRow, DataLayerError> {
let (mut query, earliest, latest) = usage_policy_window_aggregate_query(
input.windows.iter().map(|window| (window.starts_at_unix_secs, window.ends_at_unix_secs)),
"SUM(CASE WHEN state = 'finalized' THEN COALESCE(actual_cost_units, 0) ELSE reserved_cost_units END)",
)?;
query.push(" FROM usage_cost_reservations WHERE subject_id = ")
.push_bind(input.subject_id.clone())
.push(" AND admitted_at >= TO_TIMESTAMP(").push_bind(earliest)
.push("::double precision) AND admitted_at < TO_TIMESTAMP(").push_bind(latest)
.push("::double precision) AND reservation_token <> ").push_bind(input.reservation_token.clone())
.push(" AND (state = 'finalized' OR (state = 'reserved' AND reservation_expires_at > TO_TIMESTAMP(")
.push_bind(usage_policy_cost_i64(input.admitted_at_unix_secs, "usage policy admitted_at")?)
.push("::double precision)))");
query.build().fetch_one(&mut **tx).await.map_postgres_err()
}
fn settlement_from_row(
row: &sqlx::postgres::PgRow,
) -> Result<StoredUsageSettlement, DataLayerError> {
@@ -479,6 +549,8 @@ async fn consume_daily_quota_postgres(
return Ok(DailyQuotaDebitResult::default());
}
let now = chrono::Utc::now();
// Serialize each entitlement's debits. Read the shared plan's current overage policy
// from this statement's snapshot without locking every subscriber's plan row.
let entitlement_rows = sqlx::query(
r#"
SELECT
@@ -494,7 +566,7 @@ WHERE user_plan_entitlements.user_id = $1
ORDER BY user_plan_entitlements.expires_at ASC,
user_plan_entitlements.created_at ASC,
user_plan_entitlements.id ASC
FOR UPDATE
FOR UPDATE OF user_plan_entitlements
"#,
)
.bind(user_id)
@@ -655,31 +727,10 @@ WHERE event_token = $1
});
}
let window_counts = usage_policy_request_window_counts(tx, &input).await?;
for (window_index, window) in input.windows.iter().enumerate() {
let used_requests = sqlx::query_scalar::<_, i64>(
r#"
SELECT COUNT(*)::BIGINT
FROM usage_request_admissions
WHERE subject_id = $1
AND state = 'active'
AND admitted_at >= TO_TIMESTAMP($2::double precision)
AND admitted_at < TO_TIMESTAMP($3::double precision)
"#,
)
.bind(&input.subject_id)
.bind(usage_policy_cost_i64(
window.starts_at_unix_secs,
"usage policy request window start",
)?)
.bind(usage_policy_cost_i64(
window.ends_at_unix_secs,
"usage policy request window end",
)?)
.fetch_one(&mut **tx)
.await
.map_postgres_err()?;
let used_requests = usage_policy_cost_u64(
used_requests,
window_counts.try_get(window_index).map_postgres_err()?,
"usage policy request used_requests",
)?;
if used_requests >= window.limit_requests {
@@ -884,43 +935,12 @@ WHERE retain_until <= TO_TIMESTAMP($1::double precision)
.unwrap_or(0);
let target_reserved_cost_units =
previous_reserved_cost_units.max(input.reserved_cost_units);
let window_totals = usage_policy_cost_window_totals(tx, &input).await?;
for (window_index, window) in input.windows.iter().enumerate() {
let used_cost_units = sqlx::query_scalar::<_, i64>(
r#"
SELECT COALESCE(SUM(
CASE
WHEN state = 'finalized' THEN COALESCE(actual_cost_units, 0)
WHEN state = 'reserved' AND reservation_expires_at > TO_TIMESTAMP($4::double precision)
THEN reserved_cost_units
ELSE 0
END
), 0)::BIGINT
FROM usage_cost_reservations
WHERE subject_id = $1
AND admitted_at >= TO_TIMESTAMP($2::double precision)
AND admitted_at < TO_TIMESTAMP($3::double precision)
AND reservation_token <> $5
"#,
)
.bind(&input.subject_id)
.bind(usage_policy_cost_i64(
window.starts_at_unix_secs,
"usage policy window start",
)?)
.bind(usage_policy_cost_i64(
window.ends_at_unix_secs,
"usage policy window end",
)?)
.bind(usage_policy_cost_i64(
input.admitted_at_unix_secs,
"usage policy admitted_at",
)?)
.bind(&input.reservation_token)
.fetch_one(&mut **tx)
.await
.map_postgres_err()?;
let used_cost_units =
usage_policy_cost_u64(used_cost_units, "usage policy used_cost_units")?;
let used_cost_units = usage_policy_cost_u64(
window_totals.try_get(window_index).map_postgres_err()?,
"usage policy used_cost_units",
)?;
if used_cost_units
.checked_add(target_reserved_cost_units)
.is_none_or(|total| total > window.limit_cost_units)
@@ -1415,6 +1435,264 @@ WHERE id = $1
#[cfg(test)]
mod tests {
use futures_util::FutureExt;
use std::panic::AssertUnwindSafe;
async fn isolated_settlement_test_pool() -> (sqlx::PgPool, String) {
let database_url = std::env::var("AETHER_TEST_DATABASE_URL")
.expect("AETHER_TEST_DATABASE_URL must point at the test database");
let schema = format!("settlement_test_{}", uuid::Uuid::new_v4().simple());
let options = database_url
.parse::<sqlx::postgres::PgConnectOptions>()
.expect("test database URL should parse")
.options([("search_path", schema.as_str())]);
let pool = sqlx::postgres::PgPoolOptions::new()
.max_connections(2)
.connect_with(options)
.await
.expect("test database should connect");
sqlx::query(&format!("CREATE SCHEMA {schema}"))
.execute(&pool)
.await
.expect("isolated settlement schema should be created");
// Separate connections must see the same fixture, so pg_temp cannot be used here.
for table in [
"billing_plans",
"user_plan_entitlements",
"entitlement_usage_ledgers",
"users",
"usage_request_admissions",
"usage_cost_reservations",
] {
sqlx::query(&format!(
"CREATE TABLE {table} (LIKE public.{table} INCLUDING ALL)"
))
.execute(&pool)
.await
.expect("isolated settlement table should be created");
}
(pool, schema)
}
#[tokio::test]
#[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL migrations"]
async fn live_usage_policy_window_aggregates_preserve_exact_admission_and_idempotency() {
use super::*;
use aether_data_contracts::repository::settlement::{
UsagePolicyCostWindow, UsagePolicyRequestWindow,
};
let (pool, schema) = isolated_settlement_test_pool().await;
let result = AssertUnwindSafe(async {
sqlx::query("INSERT INTO users (id, username, email_verified) VALUES ('subject', 'subject', false)")
.execute(&pool).await.unwrap();
sqlx::raw_sql("INSERT INTO usage_request_admissions (request_id, subject_id, event_token, admitted_at, retain_until, state, released_at) VALUES
('older', 'subject', 'older', TO_TIMESTAMP(50), TO_TIMESTAMP(500), 'active', NULL),
('start', 'subject', 'start', TO_TIMESTAMP(100), TO_TIMESTAMP(500), 'active', NULL),
('inside', 'subject', 'inside', TO_TIMESTAMP(150), TO_TIMESTAMP(500), 'active', NULL),
('end', 'subject', 'end', TO_TIMESTAMP(200), TO_TIMESTAMP(500), 'active', NULL),
('released', 'subject', 'released', TO_TIMESTAMP(150), TO_TIMESTAMP(500), 'released', TO_TIMESTAMP(170))")
.execute(&pool).await.unwrap();
let repo = SqlxSettlementRepository::new(pool.clone());
let mut request = ReserveUsagePolicyRequestInput {
request_id: "new".to_string(), subject_id: "subject".to_string(), event_token: "new".to_string(),
admitted_at_unix_secs: 175, retain_until_unix_secs: 500,
windows: vec![
UsagePolicyRequestWindow { starts_at_unix_secs: 100, ends_at_unix_secs: 200, limit_requests: 2 },
UsagePolicyRequestWindow { starts_at_unix_secs: 0, ends_at_unix_secs: 300, limit_requests: 4 },
],
};
assert_eq!(repo.reserve_usage_policy_request(request.clone()).await.unwrap(),
ReserveUsagePolicyRequestOutcome::Rejected { window_index: 0, limit_requests: 2, used_requests: 2 });
request.windows[0].limit_requests = 3;
assert_eq!(repo.reserve_usage_policy_request(request.clone()).await.unwrap(),
ReserveUsagePolicyRequestOutcome::Rejected { window_index: 1, limit_requests: 4, used_requests: 4 });
request.windows[1].limit_requests = 5;
assert_eq!(repo.reserve_usage_policy_request(request.clone()).await.unwrap(), ReserveUsagePolicyRequestOutcome::Allowed);
assert_eq!(repo.reserve_usage_policy_request(request.clone()).await.unwrap(), ReserveUsagePolicyRequestOutcome::Allowed);
assert_eq!(sqlx::query_scalar::<_, i64>("SELECT COUNT(*) FROM usage_request_admissions WHERE event_token = 'new'")
.fetch_one(&pool).await.unwrap(), 1);
repo.release_usage_policy_request_admission(ReleaseUsagePolicyRequestAdmissionInput {
request_id: request.request_id.clone(), subject_id: request.subject_id.clone(), event_token: request.event_token.clone(), released_at_unix_secs: 180,
}).await.unwrap();
assert_eq!(repo.reserve_usage_policy_request(request).await.unwrap(), ReserveUsagePolicyRequestOutcome::AlreadyReleased);
sqlx::raw_sql("INSERT INTO usage_cost_reservations (request_id, subject_id, reservation_token, admitted_at, reserved_cost_units, actual_cost_units, state, reservation_expires_at, retain_until, finalized_at) VALUES
('older', 'subject', 'older', TO_TIMESTAMP(50), 99, 11, 'finalized', TO_TIMESTAMP(160), TO_TIMESTAMP(500), TO_TIMESTAMP(160)),
('start', 'subject', 'start', TO_TIMESTAMP(100), 99, 7, 'finalized', TO_TIMESTAMP(160), TO_TIMESTAMP(500), TO_TIMESTAMP(160)),
('inside', 'subject', 'inside', TO_TIMESTAMP(150), 5, NULL, 'reserved', TO_TIMESTAMP(300), TO_TIMESTAMP(500), NULL),
('expired', 'subject', 'expired', TO_TIMESTAMP(150), 99, NULL, 'reserved', TO_TIMESTAMP(175), TO_TIMESTAMP(500), NULL),
('end', 'subject', 'end', TO_TIMESTAMP(200), 99, 13, 'finalized', TO_TIMESTAMP(300), TO_TIMESTAMP(500), TO_TIMESTAMP(250)),
('released', 'subject', 'released', TO_TIMESTAMP(150), 99, 0, 'released', TO_TIMESTAMP(300), TO_TIMESTAMP(500), TO_TIMESTAMP(170))")
.execute(&pool).await.unwrap();
let mut cost = ReserveUsagePolicyCostInput {
request_id: "cost".to_string(), subject_id: "subject".to_string(), reservation_token: "cost".to_string(),
admitted_at_unix_secs: 175, reserved_cost_units: 3, reservation_expires_at_unix_secs: 400, retain_until_unix_secs: 500,
windows: vec![
UsagePolicyCostWindow { window_id: "short".to_string(), starts_at_unix_secs: 100, ends_at_unix_secs: 200, limit_cost_units: 14 },
UsagePolicyCostWindow { window_id: "long".to_string(), starts_at_unix_secs: 0, ends_at_unix_secs: 300, limit_cost_units: 38 },
],
};
assert_eq!(repo.reserve_usage_policy_cost(cost.clone()).await.unwrap(),
ReserveUsagePolicyCostOutcome::Rejected { window_index: 0, limit_cost_units: 14, used_cost_units: 12 });
cost.windows[0].limit_cost_units = 15;
assert_eq!(repo.reserve_usage_policy_cost(cost.clone()).await.unwrap(),
ReserveUsagePolicyCostOutcome::Rejected { window_index: 1, limit_cost_units: 38, used_cost_units: 36 });
cost.windows[1].limit_cost_units = 39;
let allowed = repo.reserve_usage_policy_cost(cost.clone()).await.unwrap();
assert!(matches!(allowed, ReserveUsagePolicyCostOutcome::Allowed { .. }), "{allowed:?}");
let repeated = repo.reserve_usage_policy_cost(cost.clone()).await.unwrap();
assert!(matches!(repeated, ReserveUsagePolicyCostOutcome::Allowed { .. }), "{repeated:?}");
cost.reserved_cost_units = 4;
assert_eq!(repo.reserve_usage_policy_cost(cost).await.unwrap(),
ReserveUsagePolicyCostOutcome::Rejected { window_index: 0, limit_cost_units: 15, used_cost_units: 12 });
sqlx::query("DELETE FROM usage_request_admissions").execute(&pool).await.unwrap();
let make_request = |id: &str| ReserveUsagePolicyRequestInput {
request_id: id.to_string(), subject_id: "subject".to_string(), event_token: id.to_string(),
admitted_at_unix_secs: 175, retain_until_unix_secs: 500,
windows: vec![UsagePolicyRequestWindow { starts_at_unix_secs: 0, ends_at_unix_secs: 300, limit_requests: 1 }],
};
let (first, second) = tokio::join!(repo.reserve_usage_policy_request(make_request("race-a")), repo.reserve_usage_policy_request(make_request("race-b")));
let outcomes = [first.unwrap(), second.unwrap()];
assert_eq!(outcomes.iter().filter(|outcome| matches!(outcome, ReserveUsagePolicyRequestOutcome::Allowed)).count(), 1);
assert_eq!(outcomes.iter().filter(|outcome| matches!(outcome, ReserveUsagePolicyRequestOutcome::Rejected { used_requests: 1, .. })).count(), 1);
}).catch_unwind().await;
sqlx::query(&format!("DROP SCHEMA {schema} CASCADE"))
.execute(&pool)
.await
.unwrap();
pool.close().await;
if let Err(panic) = result {
std::panic::resume_unwind(panic);
}
}
#[tokio::test]
#[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL migrations"]
async fn live_daily_quota_serializes_each_entitlement_without_locking_shared_plan() {
let (pool, schema) = isolated_settlement_test_pool().await;
let result = AssertUnwindSafe(async {
let grant = serde_json::json!([{
"type": "daily_quota",
"daily_quota_usd": 10.0,
"reset_timezone": "UTC",
"allow_wallet_overage": false,
}]);
sqlx::query(
"INSERT INTO billing_plans (id, title, price_amount, duration_unit, duration_value, entitlements_json, created_at, updated_at) VALUES ('shared-plan', 'Shared plan', 10, 'month', 1, $1, NOW(), NOW())",
)
.bind(&grant)
.execute(&pool)
.await
.expect("shared plan should insert");
for user_id in ["user-a", "user-b"] {
sqlx::query(
"INSERT INTO user_plan_entitlements (id, user_id, plan_id, payment_order_id, starts_at, expires_at, entitlements_snapshot, created_at, updated_at) VALUES ($1, $1, 'shared-plan', $1, NOW() - INTERVAL '1 hour', NOW() + INTERVAL '1 day', $2, NOW(), NOW())",
)
.bind(user_id)
.bind(&grant)
.execute(&pool)
.await
.expect("user entitlement should insert");
}
let mut first = pool.begin().await.expect("first transaction should start");
let first_debit = super::consume_daily_quota_postgres(
&mut first, "user-a", "request-a", 7.0, Some(0.0), false,
)
.await
.expect("first user should consume quota");
assert_eq!(first_debit.debited_usd, 7.0);
assert!(!first_debit.insufficient);
let mut second = pool.begin().await.expect("second transaction should start");
sqlx::query("SET LOCAL lock_timeout = '500ms'")
.execute(&mut *second)
.await
.expect("lock timeout should be configured");
let second_debit = super::consume_daily_quota_postgres(
&mut second, "user-b", "request-b", 2.0, Some(0.0), false,
)
.await
.expect("another user's quota must not wait for the shared plan");
assert_eq!(second_debit.debited_usd, 2.0);
assert!(!second_debit.insufficient);
second.commit().await.expect("second debit should commit");
let mut same_user = pool.begin().await.expect("contending transaction should start");
sqlx::query("SET LOCAL lock_timeout = '500ms'")
.execute(&mut *same_user)
.await
.expect("lock timeout should be configured");
let blocked = super::consume_daily_quota_postgres(
&mut same_user, "user-a", "request-a-next", 2.0, Some(0.0), false,
)
.await
.expect_err("the same entitlement must remain locked until commit");
assert!(blocked.to_string().contains("SQLSTATE 55P03"), "{blocked}");
same_user.rollback().await.expect("blocked transaction should roll back");
first.commit().await.expect("first debit should commit");
let mut next = pool.begin().await.expect("next transaction should start");
let next_debit = super::consume_daily_quota_postgres(
&mut next, "user-a", "request-a-next", 2.0, Some(0.0), false,
)
.await
.expect("same user should consume the remaining quota after commit");
assert_eq!(next_debit.debited_usd, 2.0);
assert!(!next_debit.insufficient);
next.commit().await.expect("next debit should commit");
let balance: (f64, f64) = sqlx::query_as(
"SELECT balance_before::double precision, balance_after::double precision FROM entitlement_usage_ledgers WHERE request_id = 'request-a-next'",
)
.fetch_one(&pool)
.await
.expect("next debit ledger should exist");
assert_eq!(balance, (3.0, 1.0));
let mut held = pool.begin().await.expect("quota transaction should start");
super::consume_daily_quota_postgres(
&mut held, "user-a", "request-policy-before", 0.5, Some(0.0), false,
)
.await
.expect("quota transaction should retain its entitlement lock");
let mut edit = pool.begin().await.expect("plan edit transaction should start");
sqlx::query("SET LOCAL lock_timeout = '500ms'")
.execute(&mut *edit)
.await
.expect("plan edit timeout should be configured");
sqlx::query(
"UPDATE billing_plans SET entitlements_json = jsonb_set(entitlements_json, '{0,allow_wallet_overage}', 'true'::jsonb) WHERE id = 'shared-plan'",
)
.execute(&mut *edit)
.await
.expect("plan configuration edits must not wait for usage settlement");
edit.commit().await.expect("plan edit should commit");
held.rollback().await.expect("held quota debit should roll back");
let mut after_edit = pool.begin().await.expect("fresh transaction should start");
let updated_policy = super::consume_daily_quota_postgres(
&mut after_edit, "user-a", "request-policy-after", 2.0, Some(5.0), true,
)
.await
.expect("fresh quota read should use current plan configuration");
assert!(!updated_policy.insufficient);
assert_eq!(updated_policy.debited_usd, 1.0);
after_edit.rollback().await.expect("policy verification should roll back");
})
.catch_unwind()
.await;
sqlx::query(&format!("DROP SCHEMA {schema} CASCADE"))
.execute(&pool)
.await
.expect("isolated settlement schema should be removed");
pool.close().await;
if let Err(panic) = result {
std::panic::resume_unwind(panic);
}
}
#[test]
fn finalize_usage_billing_sql_does_not_require_usage_updated_at_column() {
assert!(!super::FINALIZE_USAGE_BILLING_SQL.contains("updated_at"));
@@ -34,6 +34,14 @@ impl PostgresTransactionOptions {
}
}
pub fn maintenance() -> Self {
Self {
mode: TransactionMode::ReadWrite,
statement_timeout_ms: Some(300_000),
lock_timeout_ms: Some(30_000),
}
}
pub fn validate(&self) -> Result<(), DataLayerError> {
if matches!(self.statement_timeout_ms, Some(0)) {
return Err(DataLayerError::InvalidConfiguration(
@@ -34,7 +34,7 @@ use sqlx::{
PgPool, Postgres, QueryBuilder, Row,
};
use std::collections::{BTreeMap, BTreeSet};
use std::io::Write;
use std::io::{BufWriter, Write};
use uuid::Uuid;
use crate::{
@@ -58,6 +58,9 @@ use aether_data_contracts::repository::usage::{
use aether_data_contracts::DataLayerError;
pub mod cleanup;
mod preparation;
use preparation::prepare_usage_in_background;
// Legacy inline body columns on public.usage are deprecated. Keep the threshold at zero so
// newly captured bodies always spill to usage_body_blobs and resolve through usage_http_audits.
@@ -2121,12 +2124,9 @@ impl PreparedPendingUsage {
));
}
// Keep the capture input separate from the accounting row. The persistence sanitizer
// intentionally removes HTTP bodies/headers/states, but the pending batch still needs
// those values to populate the canonical audit/blob tables.
let capture_usage = usage.clone();
let usage = sanitize_usage_for_persistence(usage);
let prepared = prepare_usage_upsert_context(&capture_usage)?;
// Prepare captures before the accounting sanitizer removes HTTP bodies/headers/states.
let (usage, prepared) = prepare_usage_for_persistence(usage);
let prepared = prepared?;
let input_tokens = usage
.input_tokens
.map(to_i32)
@@ -8450,10 +8450,10 @@ ORDER BY "usage".user_id ASC
usage: UpsertUsageRecord,
) -> Result<StoredRequestUsageAudit, DataLayerError> {
usage.validate()?;
// `usage` is the sanitized accounting projection; prepare the auxiliary capture and
// snapshots from the original event so typed `none` markers can clear prior facts.
let capture_usage = usage.clone();
let usage = sanitize_usage_for_persistence(usage);
// Move the event before cloning or compressing captures, and do not hold a connection
// while preparing them. Stale lifecycle updates still ignore preparation errors below.
let (usage, prepared) =
prepare_usage_in_background(move || Ok(prepare_usage_for_persistence(usage))).await?;
self.tx_runner
.run_read_write(|tx| {
Box::pin(async move {
@@ -8519,7 +8519,7 @@ ORDER BY "usage".user_id ASC
clear_provider_request_body,
clear_response_body,
clear_client_response_body,
} = prepare_usage_upsert_context(&capture_usage)?;
} = prepared?;
let capture_update_allowed = recovers_terminal_failure
|| usage_capture_update_allowed(
previous_usage.as_ref().map(|stored| {
@@ -8938,33 +8938,36 @@ ORDER BY "usage".user_id ASC
return Ok(());
}
let mut request_id_counts = BTreeMap::<String, usize>::new();
for usage in &usages {
*request_id_counts
.entry(usage.request_id.clone())
.or_default() += 1;
}
// Duplicate request IDs must retain the caller's exact sequential merge order. They are
// uncommon in lifecycle batches, so keep them on the canonical single-row path.
let mut batch_rows = Vec::<(usize, PreparedPendingUsage)>::new();
let mut fallback_rows = Vec::<(usize, UpsertUsageRecord)>::new();
for (sequence, usage) in usages.into_iter().enumerate() {
let original_usage = usage.clone();
let prepared = PreparedPendingUsage::try_from_usage(usage)?;
if request_id_counts
.get(&prepared.usage.request_id)
.copied()
.unwrap_or_default()
== 1
{
batch_rows.push((sequence, prepared));
} else {
// Preserve capture markers for the canonical fallback; that path performs the
// sanitized bind only after preparing the auxiliary audit/blob state.
fallback_rows.push((sequence, original_usage));
let (batch_rows, mut fallback_rows) = prepare_usage_in_background(move || {
let mut request_id_counts = BTreeMap::<String, usize>::new();
for usage in &usages {
*request_id_counts
.entry(usage.request_id.clone())
.or_default() += 1;
}
}
// Duplicate request IDs must retain the caller's exact sequential merge order.
let mut batch_rows = Vec::<(usize, PreparedPendingUsage)>::new();
let mut fallback_rows = Vec::<(usize, UpsertUsageRecord)>::new();
for (sequence, usage) in usages.into_iter().enumerate() {
let duplicate = request_id_counts
.get(&usage.request_id)
.copied()
.unwrap_or_default()
> 1;
let original_usage = duplicate.then(|| usage.clone());
let prepared = PreparedPendingUsage::try_from_usage(usage)?;
if let Some(original_usage) = original_usage {
// Preserve capture markers for the canonical fallback, including validation
// of every row before starting the batch transaction.
fallback_rows.push((sequence, original_usage));
} else {
batch_rows.push((sequence, prepared));
}
}
Ok((batch_rows, fallback_rows))
})
.await?;
let mut inserted_request_ids = BTreeSet::<String>::new();
if !batch_rows.is_empty() {
@@ -10294,7 +10297,7 @@ RETURNING
pub async fn rebuild_api_key_usage_stats(&self) -> Result<u64, DataLayerError> {
self.tx_runner
.run_read_write(|tx| {
.run(crate::PostgresTransactionOptions::maintenance(), |tx| {
Box::pin(async move {
sqlx::query(RESET_API_KEY_USAGE_STATS_SQL)
.execute(&mut **tx)
@@ -10313,7 +10316,7 @@ RETURNING
pub async fn rebuild_provider_api_key_usage_stats(&self) -> Result<u64, DataLayerError> {
self.tx_runner
.run_read_write(|tx| {
.run(crate::PostgresTransactionOptions::maintenance(), |tx| {
Box::pin(async move {
sqlx::query(RESET_PROVIDER_API_KEY_USAGE_STATS_SQL)
.execute(&mut **tx)
@@ -12359,24 +12362,35 @@ fn prepare_usage_body_storage(value: Option<&Value>) -> Result<UsageBodyStorage,
detached_blob_bytes: None,
});
};
let bytes = serde_json::to_vec(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!("failed to serialize usage json: {err}"))
})?;
if bytes.len() == MAX_INLINE_USAGE_BODY_BYTES {
return Ok(UsageBodyStorage {
inline_json: Some(String::from_utf8(bytes).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"failed to encode inline usage body as utf-8: {err}"
))
})?),
detached_blob_bytes: None,
});
}
let mut encoder = GzEncoder::new(Vec::new(), Compression::new(6));
encoder.write_all(&bytes).map_err(|err| {
DataLayerError::UnexpectedValue(format!("failed to compress usage json: {err}"))
})?;
if MAX_INLINE_USAGE_BODY_BYTES == 0 {
// Coalesce serde's punctuation/escape writes without allocating a full JSON buffer.
let mut writer = BufWriter::with_capacity(8 * 1024, &mut encoder);
serde_json::to_writer(&mut writer, value).map_err(|err| {
let operation = if err.is_io() { "compress" } else { "serialize" };
DataLayerError::UnexpectedValue(format!("failed to {operation} usage json: {err}"))
})?;
writer.into_inner().map_err(|err| {
DataLayerError::UnexpectedValue(format!("failed to compress usage json: {err}"))
})?;
} else {
let bytes = serde_json::to_vec(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!("failed to serialize usage json: {err}"))
})?;
if bytes.len() == MAX_INLINE_USAGE_BODY_BYTES {
return Ok(UsageBodyStorage {
inline_json: Some(String::from_utf8(bytes).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"failed to encode inline usage body as utf-8: {err}"
))
})?),
detached_blob_bytes: None,
});
}
encoder.write_all(&bytes).map_err(|err| {
DataLayerError::UnexpectedValue(format!("failed to compress usage json: {err}"))
})?;
}
let detached_blob_bytes = encoder.finish().map_err(|err| {
DataLayerError::UnexpectedValue(format!("failed to finish usage json compression: {err}"))
})?;
@@ -12429,11 +12443,40 @@ fn project_usage_request_metadata(
}
}
fn prepare_usage_for_persistence(
mut usage: UpsertUsageRecord,
) -> (
UpsertUsageRecord,
Result<PreparedUsageUpsert, DataLayerError>,
) {
// Capture controls and accounting metadata have different sanitizers. Move the
// large payloads out before copying the metadata needed by both projections.
let request_body = usage.request_body.take();
let provider_request_body = usage.provider_request_body.take();
let response_body = usage.response_body.take();
let client_response_body = usage.client_response_body.take();
let request_headers = usage.request_headers.take();
let provider_request_headers = usage.provider_request_headers.take();
let response_headers = usage.response_headers.take();
let client_response_headers = usage.client_response_headers.take();
let mut capture = usage.clone();
capture.request_body = request_body;
capture.provider_request_body = provider_request_body;
capture.response_body = response_body;
capture.client_response_body = client_response_body;
capture.request_headers = request_headers;
capture.provider_request_headers = provider_request_headers;
capture.response_headers = response_headers;
capture.client_response_headers = client_response_headers;
capture.capture_retention = std::mem::take(&mut usage.capture_retention);
let capture = sanitize_usage_capture_controls_for_persistence(capture);
let prepared = prepare_usage_upsert_context(&capture);
(sanitize_usage_for_persistence(usage), prepared)
}
fn prepare_usage_upsert_context(
usage: &UpsertUsageRecord,
) -> Result<PreparedUsageUpsert, DataLayerError> {
let usage = sanitize_usage_capture_controls_for_persistence(usage.clone());
let usage = &usage;
let replace_client_request_body_facts = request_body_capture_replaces_derived_facts(
usage.request_body.as_ref(),
usage.request_body_state,
@@ -0,0 +1,327 @@
use std::sync::{Arc, OnceLock};
use std::time::Duration;
use aether_data_contracts::DataLayerError;
use tokio::sync::Semaphore;
static USAGE_PREPARATION_EXECUTOR: OnceLock<UsagePreparationExecutor> = OnceLock::new();
pub(super) async fn prepare_usage_in_background<T: Send + 'static>(
prepare: impl FnOnce() -> Result<T, DataLayerError> + Send + 'static,
) -> Result<T, DataLayerError> {
USAGE_PREPARATION_EXECUTOR
.get_or_init(|| {
UsagePreparationExecutor::new(4, 32, Duration::from_secs(1), Duration::from_secs(30))
})
.run(prepare)
.await
}
struct UsagePreparationExecutor {
workers: Arc<Semaphore>,
admitted: Arc<Semaphore>,
queue_timeout: Duration,
execution_timeout: Duration,
}
impl UsagePreparationExecutor {
fn new(
workers: usize,
admitted: usize,
queue_timeout: Duration,
execution_timeout: Duration,
) -> Self {
Self {
workers: Arc::new(Semaphore::new(workers)),
admitted: Arc::new(Semaphore::new(admitted)),
queue_timeout,
execution_timeout,
}
}
async fn run<T: Send + 'static>(
&self,
prepare: impl FnOnce() -> Result<T, DataLayerError> + Send + 'static,
) -> Result<T, DataLayerError> {
// Bound both running work and callers retaining input while waiting for a worker.
// These limits count tasks, not bytes in the caller's original usage records.
let admitted = self.admitted.clone().try_acquire_owned().map_err(|_| {
DataLayerError::TimedOut("usage preparation capacity exhausted".to_string())
})?;
let worker = tokio::time::timeout(self.queue_timeout, self.workers.clone().acquire_owned())
.await
.map_err(|_| {
DataLayerError::TimedOut(
"timed out waiting for usage preparation worker".to_string(),
)
})?
.map_err(|_| {
DataLayerError::TimedOut("usage preparation workers unavailable".to_string())
})?;
// The closure owns both permits even if its caller times out or is cancelled. It
// prepares input only; detached completion must never begin a database transaction.
let mut task = tokio::task::spawn_blocking(move || {
let _admitted = admitted;
let _worker = worker;
prepare()
});
match tokio::time::timeout(self.execution_timeout, &mut task).await {
Ok(result) => result.map_err(|error| {
DataLayerError::TimedOut(format!("usage preparation worker failed: {error}"))
})?,
Err(_) => {
// This cancels work still queued in Tokio; running blocking work keeps its
// permits until it actually exits, since abort cannot stop a blocking thread.
task.abort();
Err(DataLayerError::TimedOut(
"timed out preparing usage storage".to_string(),
))
}
}
}
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::mpsc;
use super::*;
fn executor(
admitted: usize,
queue_timeout: Duration,
execution_timeout: Duration,
) -> Arc<UsagePreparationExecutor> {
Arc::new(UsagePreparationExecutor::new(
1,
admitted,
queue_timeout,
execution_timeout,
))
}
async fn wait_for_worker_release(executor: &UsagePreparationExecutor) {
tokio::time::timeout(Duration::from_secs(2), async {
while executor.workers.available_permits() != 1
|| executor.admitted.available_permits() == 0
{
tokio::task::yield_now().await;
}
})
.await
.expect("finished blocking work should release its permits");
}
#[tokio::test(flavor = "current_thread")]
async fn preparation_runs_off_the_runtime_thread_and_preserves_errors() {
let executor = executor(1, Duration::from_secs(1), Duration::from_secs(2));
let runtime_thread = std::thread::current().id();
executor
.run(move || {
assert_ne!(std::thread::current().id(), runtime_thread);
Ok(())
})
.await
.expect("preparation should succeed");
let error = executor
.run(|| Err::<(), _>(DataLayerError::InvalidInput("bad usage".to_string())))
.await
.expect_err("input errors must reach the caller");
assert!(matches!(error, DataLayerError::InvalidInput(message) if message == "bad usage"));
}
#[tokio::test]
async fn saturated_admission_rejects_work_without_running_it() {
let executor = executor(1, Duration::from_secs(1), Duration::from_secs(2));
let (started_tx, started_rx) = tokio::sync::oneshot::channel();
let (release_tx, release_rx) = mpsc::channel();
let first_executor = executor.clone();
let first = tokio::spawn(async move {
first_executor
.run(move || {
let _ = started_tx.send(());
let _ = release_rx.recv();
Ok(())
})
.await
});
started_rx.await.expect("first job should start");
let error = executor
.run(|| -> Result<(), DataLayerError> { panic!("rejected work must not execute") })
.await
.expect_err("admission should fail immediately");
assert!(matches!(error, DataLayerError::TimedOut(message) if message.contains("capacity")));
release_tx
.send(())
.expect("first job should still be alive");
first.await.unwrap().unwrap();
}
#[tokio::test]
async fn blocking_worker_failure_is_retryable_and_releases_capacity() {
let executor = executor(1, Duration::from_secs(1), Duration::from_secs(2));
let error = executor
.run(|| -> Result<(), DataLayerError> { panic!("simulated worker failure") })
.await
.expect_err("worker failure must reach the caller");
assert!(
matches!(error, DataLayerError::TimedOut(message) if message.contains("worker failed"))
);
executor
.run(|| Ok(()))
.await
.expect("failed workers should release capacity");
}
#[tokio::test]
async fn waiting_for_a_worker_has_a_deadline_and_never_starts_expired_work() {
let executor = executor(2, Duration::from_millis(20), Duration::from_secs(2));
let (started_tx, started_rx) = tokio::sync::oneshot::channel();
let (release_tx, release_rx) = mpsc::channel();
let first_executor = executor.clone();
let first = tokio::spawn(async move {
first_executor
.run(move || {
let _ = started_tx.send(());
let _ = release_rx.recv();
Ok(())
})
.await
});
started_rx.await.expect("first job should start");
let ran = Arc::new(AtomicBool::new(false));
let work_ran = ran.clone();
let error = executor
.run(move || {
work_ran.store(true, Ordering::SeqCst);
Ok(())
})
.await
.expect_err("the queued job should time out");
assert!(matches!(error, DataLayerError::TimedOut(message) if message.contains("waiting")));
assert_eq!(executor.admitted.available_permits(), 1);
release_tx
.send(())
.expect("first job should still be alive");
first.await.unwrap().unwrap();
assert!(!ran.load(Ordering::SeqCst));
}
#[tokio::test]
async fn cancellation_keeps_permits_until_running_blocking_work_exits() {
let executor = executor(1, Duration::from_secs(1), Duration::from_secs(2));
let (started_tx, started_rx) = tokio::sync::oneshot::channel();
let (release_tx, release_rx) = mpsc::channel();
let first_executor = executor.clone();
let first = tokio::spawn(async move {
first_executor
.run(move || {
let _ = started_tx.send(());
let _ = release_rx.recv();
Ok(())
})
.await
});
started_rx.await.expect("first job should start");
first.abort();
assert!(first.await.unwrap_err().is_cancelled());
assert_eq!(executor.workers.available_permits(), 0);
assert!(matches!(
executor.run(|| Ok(())).await,
Err(DataLayerError::TimedOut(_))
));
release_tx
.send(())
.expect("blocking work should outlive cancellation");
wait_for_worker_release(&executor).await;
executor
.run(|| Ok(()))
.await
.expect("the executor should recover");
}
#[tokio::test]
async fn cancellation_keeps_capture_budget_until_blocking_input_is_dropped() {
use aether_data_contracts::repository::usage::{
usage_json_heap_estimate, UpsertUsageRecord, UsageCaptureMemoryBudget,
};
let mut usage: UpsertUsageRecord = serde_json::from_value(serde_json::json!({
"request_id": "req-cancelled-preparation",
"provider_name": "test",
"model": "test",
"status": "completed",
"billing_status": "pending",
"updated_at_unix_secs": 100,
"request_body": {"content": "retained".repeat(1024)}
}))
.unwrap();
let bytes = std::mem::size_of::<serde_json::Value>()
+ usage_json_heap_estimate(usage.request_body.as_ref().unwrap());
let budget = Arc::new(UsageCaptureMemoryBudget::new(bytes));
assert!(usage.capture_retention.reserve(Arc::clone(&budget), bytes));
let executor = executor(1, Duration::from_secs(1), Duration::from_secs(2));
let (started_tx, started_rx) = tokio::sync::oneshot::channel();
let (release_tx, release_rx) = mpsc::channel();
let first_executor = Arc::clone(&executor);
let first = tokio::spawn(async move {
first_executor
.run(move || {
let _ = started_tx.send(());
let _ = release_rx.recv();
drop(usage);
Ok(())
})
.await
});
started_rx.await.unwrap();
first.abort();
assert!(first.await.unwrap_err().is_cancelled());
assert_eq!(budget.retained_bytes(), bytes);
release_tx.send(()).unwrap();
wait_for_worker_release(&executor).await;
assert_eq!(budget.retained_bytes(), 0);
}
#[tokio::test]
async fn execution_timeout_keeps_permits_until_running_blocking_work_exits() {
let executor = executor(1, Duration::from_secs(1), Duration::from_millis(20));
let (started_tx, started_rx) = tokio::sync::oneshot::channel();
let (release_tx, release_rx) = mpsc::channel();
let first_executor = executor.clone();
let first = tokio::spawn(async move {
first_executor
.run(move || {
let _ = started_tx.send(());
let _ = release_rx.recv();
Ok(())
})
.await
});
started_rx.await.expect("first job should start");
let error = first
.await
.unwrap()
.expect_err("running work should time out");
assert!(
matches!(error, DataLayerError::TimedOut(message) if message.contains("preparing"))
);
assert_eq!(executor.workers.available_permits(), 0);
assert!(matches!(
executor.run(|| Ok(())).await,
Err(DataLayerError::TimedOut(_))
));
release_tx
.send(())
.expect("blocking work should outlive timeout");
wait_for_worker_release(&executor).await;
executor
.run(|| Ok(()))
.await
.expect("the executor should recover");
}
}
@@ -8,7 +8,7 @@ use super::{
attach_usage_routing_snapshot_metadata, attach_usage_settlement_pricing_snapshot_metadata,
clear_previous_request_body_facts, inflate_usage_json_value,
prepare_request_metadata_for_body_storage, prepare_usage_body_storage,
prepare_usage_upsert_context, push_postgres_usage_websocket_filter,
prepare_usage_for_persistence, push_postgres_usage_websocket_filter,
request_body_capture_replaces_derived_facts, resolved_read_usage_body_ref,
resolved_write_usage_body_ref, split_dashboard_daily_aggregate_range,
split_dashboard_hourly_aggregate_range, usage_body_capture_state_for_storage, usage_body_ref,
@@ -39,6 +39,7 @@ fn fast_clear_usage_record(
terminal_service_tier: Option<&str>,
) -> UpsertUsageRecord {
UpsertUsageRecord {
capture_retention: Default::default(),
request_id: request_id.to_string(),
user_id: None,
api_key_id: None,
@@ -204,7 +205,26 @@ async fn live_full_http_capture_round_trips_for_direct_and_batch_writes() {
let repository = SqlxUsageReadRepository::new(factory.connect_lazy().unwrap());
crate::run_migrations(repository.pool()).await.unwrap();
for batch in [false, true] {
for write_mode in 0..3 {
use aether_data_contracts::repository::usage::{
usage_json_heap_estimate, UsageCaptureMemoryBudget,
};
let batch = write_mode != 0;
let budget = Arc::new(UsageCaptureMemoryBudget::new(4 * 1024 * 1024));
let retain_capture = |usage: &mut UpsertUsageRecord| {
let bytes = [
usage.request_body.as_ref(),
usage.provider_request_body.as_ref(),
usage.response_body.as_ref(),
usage.client_response_body.as_ref(),
]
.into_iter()
.flatten()
.map(|body| std::mem::size_of::<serde_json::Value>() + usage_json_heap_estimate(body))
.sum();
assert!(usage.capture_retention.reserve(Arc::clone(&budget), bytes));
bytes
};
let request_id = format!("req-full-capture-{}", uuid::Uuid::new_v4().simple());
let now_unix_secs = Utc::now().timestamp() as u64;
let mut pending = fast_clear_usage_record(
@@ -225,14 +245,18 @@ async fn live_full_http_capture_round_trips_for_direct_and_batch_writes() {
pending.response_body_state = Some(UsageBodyCaptureState::Inline);
pending.client_response_body = Some(json!("pending client response"));
pending.client_response_body_state = Some(UsageBodyCaptureState::Inline);
let pending_bytes = retain_capture(&mut pending);
if batch {
repository
.upsert_pending_many(vec![pending.clone()])
.await
.unwrap();
let records = if write_mode == 2 {
vec![pending.clone(), pending.clone()]
} else {
vec![pending.clone()]
};
repository.upsert_pending_many(records).await.unwrap();
} else {
repository.upsert(pending.clone()).await.unwrap();
}
assert_eq!(budget.retained_bytes(), pending_bytes);
for (field, expected) in [
(UsageBodyField::RequestBody, pending.request_body.as_ref()),
(
@@ -274,7 +298,10 @@ async fn live_full_http_capture_round_trips_for_direct_and_batch_writes() {
terminal.response_body_state = Some(UsageBodyCaptureState::Inline);
terminal.client_response_body = Some(json!({"output": "final response"}));
terminal.client_response_body_state = Some(UsageBodyCaptureState::Inline);
let terminal_bytes = retain_capture(&mut terminal);
repository.upsert(terminal.clone()).await.unwrap();
assert_eq!(budget.retained_bytes(), pending_bytes + terminal_bytes);
assert_eq!(budget.downgraded_total(), 0);
let stored = repository
.find_by_request_id_shallow(&request_id)
@@ -330,6 +357,9 @@ async fn live_full_http_capture_round_trips_for_direct_and_batch_writes() {
.execute(repository.pool())
.await
.unwrap();
drop(pending);
drop(terminal);
assert_eq!(budget.retained_bytes(), 0);
}
}
@@ -2411,6 +2441,7 @@ async fn validates_upsert_before_hitting_database() {
let repository = SqlxUsageReadRepository::new(pool);
let result = repository
.upsert(UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "".to_string(),
user_id: None,
api_key_id: None,
@@ -4324,6 +4355,96 @@ fn prepare_usage_body_storage_compresses_large_payloads() {
);
}
#[test]
fn prepare_usage_body_storage_streams_json_shapes_into_compatible_gzip() {
for payload in [
serde_json::Value::Null,
json!(false),
json!(42),
json!(["quoted\"text", "line\nbreak", "\u{4e2d}\u{6587}", null]),
json!({
"content": "escaped\n\"\\value".repeat(32 * 1024),
"nested": {"values": [true, null, 1.25, -7]}
}),
] {
let storage = prepare_usage_body_storage(Some(&payload)).expect("body should compress");
assert!(storage.inline_json.is_none());
let compressed = storage
.detached_blob_bytes
.expect("body should be detached");
assert_eq!(
inflate_usage_json_value(&compressed).expect("body should remain readable"),
payload
);
}
}
#[test]
fn managed_capture_preparation_moves_bodies_without_a_second_reservation() {
use aether_data_contracts::repository::usage::{
sanitize_usage_for_persistence, usage_json_heap_estimate, UsageCaptureMemoryBudget,
};
let mut usage = fast_clear_usage_record(
"req-managed-capture",
"managed-capture",
100,
true,
UsageBodyCaptureState::Inline,
Some("priority"),
);
let bodies = [
json!({"messages": [{"role": "user", "content": "request".repeat(4096)}]}),
json!({"input": "provider request".repeat(4096), "service_tier": "priority"}),
json!({"output": "provider response".repeat(4096)}),
json!({"output": "client response".repeat(4096)}),
];
usage.request_body = Some(bodies[0].clone());
usage.provider_request_body = Some(bodies[1].clone());
usage.response_body = Some(bodies[2].clone());
usage.client_response_body = Some(bodies[3].clone());
usage.request_body_state = Some(UsageBodyCaptureState::Inline);
usage.response_body_state = Some(UsageBodyCaptureState::Inline);
usage.client_response_body_state = Some(UsageBodyCaptureState::Inline);
usage.request_headers = Some(json!({"content-type": "application/json"}));
usage.cache_read_input_tokens = Some(0);
usage.total_cost_usd = Some(0.25);
usage.actual_total_cost_usd = Some(0.125);
let expected_accounting = sanitize_usage_for_persistence(usage.clone());
let bytes = [
usage.request_body.as_ref(),
usage.provider_request_body.as_ref(),
usage.response_body.as_ref(),
usage.client_response_body.as_ref(),
]
.into_iter()
.flatten()
.map(|body| std::mem::size_of::<serde_json::Value>() + usage_json_heap_estimate(body))
.sum();
let budget = Arc::new(UsageCaptureMemoryBudget::new(bytes));
assert!(usage.capture_retention.reserve(Arc::clone(&budget), bytes));
let (accounting, prepared) = prepare_usage_for_persistence(usage);
let prepared = prepared.expect("managed capture should prepare without cloning bodies");
assert_eq!(budget.retained_bytes(), 0);
assert_eq!(budget.downgraded_total(), 0);
assert_eq!(accounting, expected_accounting);
for (storage, expected) in [
prepared.request_body_storage,
prepared.provider_request_body_storage,
prepared.response_body_storage,
prepared.client_response_body_storage,
]
.into_iter()
.zip(bodies)
{
assert_eq!(
inflate_usage_json_value(storage.detached_blob_bytes.as_deref().unwrap()).unwrap(),
expected
);
}
}
#[test]
fn usage_body_capture_state_for_storage_marks_detached_bodies_as_reference() {
let payload = json!({"message": "hello"});
@@ -4413,7 +4534,8 @@ fn explicit_none_capture_drops_residual_body_ref_and_incoming_fast_metadata_befo
"provider_request_body_ref": "usage://request/req-none-residual/provider_request_body"
}));
let prepared = prepare_usage_upsert_context(&usage).expect("usage should prepare");
let (_, prepared) = prepare_usage_for_persistence(usage);
let prepared = prepared.expect("usage should prepare");
assert!(prepared.clear_provider_request_body);
assert!(!prepared.provider_request_body_storage.has_detached_blob());
assert_eq!(prepared.http_audit_refs.provider_request_body_ref, None);
@@ -4805,6 +4927,7 @@ fn attach_usage_http_audit_body_refs_adds_missing_metadata_without_overwriting_e
fn usage_routing_snapshot_from_usage_only_activates_for_routing_metadata() {
let snapshot = usage_routing_snapshot_from_usage(
&UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-123".to_string(),
user_id: None,
api_key_id: None,
@@ -4904,6 +5027,7 @@ fn usage_routing_snapshot_from_usage_only_activates_for_routing_metadata() {
let empty_snapshot = usage_routing_snapshot_from_usage(
&UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-124".to_string(),
user_id: None,
api_key_id: None,
@@ -4982,6 +5106,7 @@ fn usage_routing_snapshot_from_usage_only_activates_for_routing_metadata() {
fn usage_routing_snapshot_from_usage_prefers_typed_routing_fields_without_metadata() {
let snapshot = usage_routing_snapshot_from_usage(
&UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-typed-routing-1".to_string(),
user_id: None,
api_key_id: None,
@@ -5116,6 +5241,7 @@ fn attach_usage_routing_snapshot_metadata_adds_missing_keys_without_overwriting_
fn usage_settlement_pricing_snapshot_from_usage_extracts_typed_billing_fields() {
let snapshot = usage_settlement_pricing_snapshot_from_usage(
&UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-125".to_string(),
user_id: None,
api_key_id: None,
@@ -224,6 +224,36 @@ pub struct StoredRequestCandidate {
}
impl StoredRequestCandidate {
/// Scheduling needs identity, status, counters and times, without diagnostic payloads.
pub fn runtime_snapshot(&self) -> Self {
Self {
id: self.id.clone(),
request_id: self.request_id.clone(),
user_id: self.user_id.clone(),
api_key_id: self.api_key_id.clone(),
username: None,
api_key_name: None,
candidate_index: self.candidate_index,
retry_index: self.retry_index,
provider_id: self.provider_id.clone(),
endpoint_id: self.endpoint_id.clone(),
key_id: self.key_id.clone(),
status: self.status,
skip_reason: None,
is_cached: self.is_cached,
status_code: self.status_code,
error_type: None,
error_message: None,
latency_ms: self.latency_ms,
concurrent_requests: self.concurrent_requests,
extra_data: None,
required_capabilities: None,
created_at_unix_ms: self.created_at_unix_ms,
started_at_unix_ms: self.started_at_unix_ms,
finished_at_unix_ms: self.finished_at_unix_ms,
}
}
pub fn sanitize_for_persistence(&mut self) {
self.username = None;
self.api_key_name = None;
@@ -691,6 +721,19 @@ pub trait RequestCandidateReadRepository: Send + Sync {
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, crate::DataLayerError>;
/// Same ordering and limit as `list_recent`, omitting diagnostic fields.
async fn list_recent_runtime(
&self,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, crate::DataLayerError> {
Ok(self
.list_recent(limit)
.await?
.iter()
.map(StoredRequestCandidate::runtime_snapshot)
.collect())
}
async fn list_by_provider_id(
&self,
provider_id: &str,
@@ -849,9 +892,7 @@ pub fn sanitize_request_candidate_extra_data_for_persistence(
extra_data: Option<serde_json::Value>,
) -> Option<serde_json::Value> {
let object = extra_data.as_ref()?.as_object()?;
let mut sanitized = sanitize_request_candidate_extra_data(extra_data.clone())
.and_then(|value| value.as_object().cloned())
.unwrap_or_default();
let mut sanitized = sanitize_candidate_extra_data_object(object);
for (key, fields) in [
("upstream_response", &["headers", "body"][..]),
("error_flow", &["message"][..]),
@@ -884,7 +925,10 @@ pub fn sanitize_request_candidate_extra_data_for_persistence(
};
let mut summary = sanitized
.remove(key)
.and_then(|value| value.as_object().cloned())
.and_then(|value| match value {
serde_json::Value::Object(object) => Some(object),
_ => None,
})
.unwrap_or_default();
for field in fields {
if let Some(value) = diagnostic.get(*field).filter(|value| !value.is_null()) {
@@ -907,70 +951,72 @@ pub fn sanitize_request_candidate_extra_data(
let serde_json::Value::Object(object) = extra_data? else {
return None;
};
let sanitized = sanitize_candidate_extra_data_object(&object);
(!sanitized.is_empty()).then_some(serde_json::Value::Object(sanitized))
}
fn sanitize_candidate_extra_data_object(
object: &serde_json::Map<String, serde_json::Value>,
) -> serde_json::Map<String, serde_json::Value> {
let mut sanitized = serde_json::Map::new();
for field in ["gateway_execution_runtime", "stream_completed", "cache_1h"] {
insert_candidate_bool(&object, &mut sanitized, field);
insert_candidate_bool(object, &mut sanitized, field);
}
for field in ["first_byte_time_ms", "pool_key_index"] {
insert_candidate_u64(&object, &mut sanitized, field);
insert_candidate_u64(object, &mut sanitized, field);
}
insert_candidate_i64(&object, &mut sanitized, "priority_slot");
insert_candidate_u64(&object, &mut sanitized, "ranking_index");
insert_candidate_i64(object, &mut sanitized, "priority_slot");
insert_candidate_u64(object, &mut sanitized, "ranking_index");
insert_candidate_known_string(&object, &mut sanitized, "phase", sanitize_candidate_phase);
insert_candidate_known_string(object, &mut sanitized, "phase", sanitize_candidate_phase);
for field in [
"client_api_format",
"provider_api_format",
"client_contract",
"provider_contract",
] {
insert_candidate_known_string(
&object,
&mut sanitized,
field,
sanitize_candidate_api_format,
);
insert_candidate_known_string(object, &mut sanitized, field, sanitize_candidate_api_format);
}
insert_candidate_known_string(
&object,
object,
&mut sanitized,
"execution_strategy",
sanitize_candidate_execution_strategy,
);
insert_candidate_known_string(
&object,
object,
&mut sanitized,
"conversion_mode",
sanitize_candidate_conversion_mode,
);
insert_candidate_known_string(
&object,
object,
&mut sanitized,
"ranking_mode",
sanitize_candidate_ranking_mode,
);
insert_candidate_known_string(
&object,
object,
&mut sanitized,
"priority_mode",
sanitize_candidate_priority_mode,
);
insert_candidate_known_string(
&object,
object,
&mut sanitized,
"promoted_by",
sanitize_candidate_promotion_reason,
);
insert_candidate_known_string(
&object,
object,
&mut sanitized,
"demoted_by",
sanitize_candidate_demotion_reason,
);
insert_candidate_known_string(&object, &mut sanitized, "source", sanitize_candidate_source);
insert_candidate_known_string(object, &mut sanitized, "source", sanitize_candidate_source);
insert_candidate_known_string(
&object,
object,
&mut sanitized,
"execution_path",
sanitize_candidate_execution_path,
@@ -1030,7 +1076,7 @@ pub fn sanitize_request_candidate_extra_data(
sanitized.insert("pool_group_exhaustion".to_string(), exhaustion);
}
(!sanitized.is_empty()).then_some(serde_json::Value::Object(sanitized))
sanitized
}
pub fn sanitize_request_candidate_required_capabilities(
@@ -0,0 +1,490 @@
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use serde_json::{Map, Value};
use super::types::UsageBodyCaptureState;
/// Shared accounting for the estimated heap retained by diagnostic JSON bodies.
#[doc(hidden)]
#[derive(Debug)]
pub struct UsageCaptureMemoryBudget {
limit: usize,
retained: AtomicUsize,
downgraded_total: AtomicU64,
}
impl UsageCaptureMemoryBudget {
pub fn new(limit: usize) -> Self {
Self {
limit,
retained: AtomicUsize::new(0),
downgraded_total: AtomicU64::new(0),
}
}
fn try_reserve(&self, bytes: usize) -> bool {
if bytes == 0 {
return true;
}
self.retained
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |retained| {
retained
.checked_add(bytes)
.filter(|next| *next <= self.limit)
})
.is_ok()
}
fn release(&self, bytes: usize) {
if bytes != 0 {
self.retained.fetch_sub(bytes, Ordering::AcqRel);
}
}
fn record_downgrade(&self) {
self.downgraded_total.fetch_add(1, Ordering::Relaxed);
}
pub fn retained_bytes(&self) -> usize {
self.retained.load(Ordering::Acquire)
}
pub fn downgraded_total(&self) -> u64 {
self.downgraded_total.load(Ordering::Relaxed)
}
pub fn snapshot(&self) -> (usize, usize, u64) {
(self.limit, self.retained_bytes(), self.downgraded_total())
}
}
/// Non-serialized ownership of the diagnostic JSON heap estimate.
#[doc(hidden)]
#[derive(Debug, Default)]
pub struct UsageCaptureRetention {
budget: Option<Arc<UsageCaptureMemoryBudget>>,
bytes: usize,
}
// Runtime accounting does not participate in value or wire equality.
impl PartialEq for UsageCaptureRetention {
fn eq(&self, _other: &Self) -> bool {
true
}
}
impl UsageCaptureRetention {
pub fn reserve(&mut self, budget: Arc<UsageCaptureMemoryBudget>, bytes: usize) -> bool {
if self
.budget
.as_ref()
.is_some_and(|current| Arc::ptr_eq(current, &budget))
{
if bytes > self.bytes && !budget.try_reserve(bytes - self.bytes) {
budget.record_downgrade();
return false;
}
if bytes < self.bytes {
budget.release(self.bytes - bytes);
}
self.bytes = bytes;
return true;
}
if !budget.try_reserve(bytes) {
budget.record_downgrade();
return false;
}
*self = Self {
budget: Some(budget),
bytes,
};
true
}
pub fn clear(&mut self, budget: Arc<UsageCaptureMemoryBudget>) {
*self = Self {
budget: Some(budget),
bytes: 0,
};
}
pub fn clone_for_bodies(&self, estimate: impl FnOnce() -> usize) -> (Self, bool) {
let Some(budget) = &self.budget else {
return (Self::default(), true);
};
let mut retention = Self::default();
if retention.reserve(Arc::clone(budget), estimate()) {
(retention, true)
} else {
retention.clear(Arc::clone(budget));
(retention, false)
}
}
}
impl Drop for UsageCaptureRetention {
fn drop(&mut self) {
if let Some(budget) = &self.budget {
budget.release(self.bytes);
}
}
}
// serde_json::Map does not expose its backing allocation capacity. This charges a
// conservative per-entry estimate, not an allocator or process RSS measurement.
#[doc(hidden)]
pub fn usage_json_heap_estimate(value: &Value) -> usize {
match value {
Value::String(value) => value.capacity(),
Value::Array(values) => values.iter().fold(
values
.capacity()
.saturating_mul(std::mem::size_of::<Value>()),
|bytes, value| bytes.saturating_add(usage_json_heap_estimate(value)),
),
Value::Object(values) => values.iter().fold(
values.len().saturating_mul(
4 * (std::mem::size_of::<String>()
+ std::mem::size_of::<Value>()
+ std::mem::size_of::<usize>()),
),
|bytes, (key, value)| {
bytes
.saturating_add(key.capacity())
.saturating_add(usage_json_heap_estimate(value))
},
),
Value::Null | Value::Bool(_) | Value::Number(_) => 0,
}
}
/// Marks an omitted diagnostic body using the existing capture metadata shape.
#[doc(hidden)]
pub fn mark_usage_capture_memory_omitted(metadata: &mut Option<Value>, key: &str) {
let source_bytes = metadata
.as_ref()
.and_then(|metadata| metadata.get("body_capture"))
.and_then(|capture| capture.get(key))
.and_then(|entry| entry.get("source_bytes"))
.and_then(Value::as_u64);
let Some(metadata) = metadata
.get_or_insert_with(|| Value::Object(Map::with_capacity(1)))
.as_object_mut()
else {
return;
};
let Some(body_capture) = metadata
.entry("body_capture".to_owned())
.or_insert_with(|| Value::Object(Map::with_capacity(1)))
.as_object_mut()
else {
return;
};
let mut entry = Map::with_capacity(3 + usize::from(source_bytes.is_some()));
entry.insert(
"state".to_owned(),
Value::String(UsageBodyCaptureState::Truncated.as_str().to_owned()),
);
entry.insert("stored_bytes".to_owned(), Value::from(0));
if let Some(source_bytes) = source_bytes {
entry.insert("source_bytes".to_owned(), Value::from(source_bytes));
}
entry.insert(
"reason".to_owned(),
Value::String("usage_event_memory_budget_exceeded".to_owned()),
);
body_capture.insert(key.to_owned(), Value::Object(entry));
}
#[cfg(test)]
mod tests {
use super::super::types::UpsertUsageRecord;
use super::*;
use serde_json::json;
fn upsert_with_diagnostic_bodies() -> UpsertUsageRecord {
serde_json::from_str(
r#"{
"request_id": "retained-request",
"provider_name": "openai",
"model": "test-model",
"status": "completed",
"billing_status": "settled",
"updated_at_unix_secs": 123,
"input_tokens": 100,
"output_tokens": 500,
"total_tokens": 600,
"cache_creation_input_tokens": 7,
"cache_creation_ephemeral_5m_input_tokens": 7,
"cache_creation_ephemeral_1h_input_tokens": 0,
"cache_read_input_tokens": 0,
"total_cost_usd": 1.25,
"actual_total_cost_usd": 0.75,
"cache_creation_cost_usd": 0.05,
"cache_read_cost_usd": 0.0,
"status_code": 200,
"error_message": "preserved diagnostic classification",
"request_headers": {"x-request": "original"},
"provider_request_headers": {"x-provider-request": "original"},
"response_headers": {"x-response": "original"},
"client_response_headers": {"x-client-response": "original"},
"request_body": {"text": "original request body"},
"provider_request_body": {"text": "original provider request body"},
"response_body": {"text": "original provider response body"},
"client_response_body": {"text": "original client response body"},
"request_body_ref": "usage://retained-request/request",
"provider_request_body_ref": "usage://retained-request/provider_request",
"response_body_ref": "usage://retained-request/response",
"client_response_body_ref": "usage://retained-request/client_response",
"request_body_state": "inline",
"provider_request_body_state": "reference",
"response_body_state": "inline",
"client_response_body_state": "truncated",
"request_metadata": {
"trace_id": "unchanged",
"body_capture": {
"request": {"state": "inline", "source_bytes": 100},
"provider_request": {"state": "reference", "source_bytes": 200},
"response": {"state": "inline", "source_bytes": 300},
"client_response": {"state": "truncated", "source_bytes": 400}
}
}
}"#,
)
.expect("valid usage write fixture")
}
fn upsert_body_estimate(record: &UpsertUsageRecord) -> usize {
[
&record.request_body,
&record.provider_request_body,
&record.response_body,
&record.client_response_body,
]
.into_iter()
.flatten()
.map(|body| std::mem::size_of::<Value>() + usage_json_heap_estimate(body))
.sum()
}
#[test]
fn usage_capture_memory_upsert_serde_skips_retention_and_preserves_value_equality() {
let mut source = upsert_with_diagnostic_bodies();
let weight = upsert_body_estimate(&source);
let budget = Arc::new(UsageCaptureMemoryBudget::new(weight));
assert!(source
.capture_retention
.reserve(Arc::clone(&budget), weight));
let mut serialized = serde_json::to_value(&source).unwrap();
assert!(serialized.get("capture_retention").is_none());
serialized["capture_retention"] = json!({"bytes": usize::MAX});
let roundtrip: UpsertUsageRecord = serde_json::from_value(serialized).unwrap();
assert_eq!(source, roundtrip);
let unmanaged_clone = roundtrip.clone();
assert_eq!(source, unmanaged_clone);
assert!(unmanaged_clone.request_body.is_some());
assert!(unmanaged_clone.provider_request_body.is_some());
assert!(unmanaged_clone.response_body.is_some());
assert!(unmanaged_clone.client_response_body.is_some());
assert_eq!(budget.retained_bytes(), weight);
assert_eq!(budget.downgraded_total(), 0);
drop((roundtrip, unmanaged_clone));
assert_eq!(budget.retained_bytes(), weight);
drop(source);
assert_eq!(budget.retained_bytes(), 0);
}
#[test]
fn usage_capture_memory_upsert_clone_reserves_for_all_four_deep_copies() {
let mut source = upsert_with_diagnostic_bodies();
let original = serde_json::to_value(&source).unwrap();
let weight = upsert_body_estimate(&source);
let budget = Arc::new(UsageCaptureMemoryBudget::new(weight * 2));
assert!(source
.capture_retention
.reserve(Arc::clone(&budget), weight));
let cloned = source.clone();
assert_eq!(source, cloned);
assert_eq!(budget.retained_bytes(), weight * 2);
assert_eq!(budget.downgraded_total(), 0);
for (source_body, cloned_body) in [
(&source.request_body, &cloned.request_body),
(&source.provider_request_body, &cloned.provider_request_body),
(&source.response_body, &cloned.response_body),
(&source.client_response_body, &cloned.client_response_body),
] {
let source_text = source_body.as_ref().unwrap()["text"].as_str().unwrap();
let cloned_text = cloned_body.as_ref().unwrap()["text"].as_str().unwrap();
assert_eq!(source_text, cloned_text);
assert_ne!(source_text.as_ptr(), cloned_text.as_ptr());
}
assert_eq!(serde_json::to_value(&source).unwrap(), original);
drop(source);
assert_eq!(budget.retained_bytes(), weight);
assert_eq!(serde_json::to_value(&cloned).unwrap(), original);
drop(cloned);
assert_eq!(budget.retained_bytes(), 0);
}
#[test]
fn usage_capture_memory_upsert_clone_over_budget_only_omits_four_bodies() {
let mut source = upsert_with_diagnostic_bodies();
let original = serde_json::to_value(&source).unwrap();
let weight = upsert_body_estimate(&source);
let budget = Arc::new(UsageCaptureMemoryBudget::new(weight));
assert!(source
.capture_retention
.reserve(Arc::clone(&budget), weight));
let cloned = source.clone();
assert_eq!(budget.retained_bytes(), weight);
assert_eq!(budget.downgraded_total(), 1);
assert_eq!(serde_json::to_value(&source).unwrap(), original);
let mut expected = original;
for (body, state, key, source_bytes) in [
("request_body", "request_body_state", "request", 100),
(
"provider_request_body",
"provider_request_body_state",
"provider_request",
200,
),
("response_body", "response_body_state", "response", 300),
(
"client_response_body",
"client_response_body_state",
"client_response",
400,
),
] {
expected[body] = Value::Null;
expected[state] = json!("truncated");
expected["request_metadata"]["body_capture"][key] = json!({
"state": "truncated",
"stored_bytes": 0,
"source_bytes": source_bytes,
"reason": "usage_event_memory_budget_exceeded"
});
}
assert_eq!(serde_json::to_value(&cloned).unwrap(), expected);
assert_eq!(cloned.output_tokens, Some(500));
assert_eq!(cloned.cache_read_input_tokens, Some(0));
assert_eq!(cloned.cache_creation_ephemeral_1h_input_tokens, Some(0));
assert_eq!(cloned.total_cost_usd, Some(1.25));
assert_eq!(cloned.actual_total_cost_usd, Some(0.75));
drop(cloned);
assert_eq!(budget.retained_bytes(), weight);
drop(source);
assert_eq!(budget.retained_bytes(), 0);
}
#[test]
fn usage_capture_memory_upsert_clone_preserves_explicit_clearing_states() {
for state in [
UsageBodyCaptureState::None,
UsageBodyCaptureState::Disabled,
UsageBodyCaptureState::Unavailable,
] {
let mut source = upsert_with_diagnostic_bodies();
source.request_body_state = Some(state);
source.provider_request_body_state = Some(state);
source.response_body_state = Some(state);
source.client_response_body_state = Some(state);
let original = serde_json::to_value(&source).unwrap();
let weight = upsert_body_estimate(&source);
let budget = Arc::new(UsageCaptureMemoryBudget::new(weight));
assert!(source
.capture_retention
.reserve(Arc::clone(&budget), weight));
let cloned = source.clone();
let mut expected = original.clone();
for body in [
"request_body",
"provider_request_body",
"response_body",
"client_response_body",
] {
expected[body] = Value::Null;
}
assert_eq!(serde_json::to_value(&cloned).unwrap(), expected);
assert_eq!(serde_json::to_value(&source).unwrap(), original);
assert_eq!(budget.retained_bytes(), weight);
assert_eq!(budget.downgraded_total(), 1);
drop(cloned);
assert_eq!(budget.retained_bytes(), weight);
drop(source);
assert_eq!(budget.retained_bytes(), 0);
}
}
#[test]
fn usage_capture_memory_metadata_preserves_source_bytes_and_unrelated_metadata() {
let mut metadata = Some(json!({
"trace_id": "unchanged",
"body_capture": {
"request": {"state": "complete", "source_bytes": 42, "stored_bytes": 42, "extra": true},
"response": {"state": "complete", "source_bytes": 7}
}
}));
mark_usage_capture_memory_omitted(&mut metadata, "request");
assert_eq!(
metadata,
Some(json!({
"trace_id": "unchanged",
"body_capture": {
"request": {
"state": "truncated",
"source_bytes": 42,
"stored_bytes": 0,
"reason": "usage_event_memory_budget_exceeded"
},
"response": {"state": "complete", "source_bytes": 7}
}
}))
);
}
#[test]
fn usage_capture_memory_metadata_creates_missing_objects_and_replaces_entries() {
for mut metadata in [
None,
Some(json!({})),
Some(json!({"body_capture": {}})),
Some(json!({"body_capture": {"request": null}})),
Some(json!({"body_capture": {"request": "legacy"}})),
Some(json!({"body_capture": {"request": {"source_bytes": "42"}}})),
Some(json!({"body_capture": {"request": {"source_bytes": -1}}})),
] {
mark_usage_capture_memory_omitted(&mut metadata, "request");
assert_eq!(
metadata,
Some(json!({"body_capture": {"request": {
"state": "truncated",
"stored_bytes": 0,
"reason": "usage_event_memory_budget_exceeded"
}}}))
);
}
}
#[test]
fn usage_capture_memory_metadata_preserves_existing_non_object_containers() {
for metadata in [
Value::Null,
json!(false),
json!(7),
json!("legacy"),
json!([]),
json!({"body_capture": null}),
json!({"body_capture": false}),
json!({"body_capture": 7}),
json!({"body_capture": "legacy"}),
json!({"body_capture": []}),
] {
let mut actual = Some(metadata.clone());
mark_usage_capture_memory_omitted(&mut actual, "request");
assert_eq!(actual, Some(metadata));
}
}
}
@@ -1,8 +1,14 @@
mod capture_memory;
mod compression;
mod metadata_policy;
mod policy;
mod types;
#[doc(hidden)]
pub use capture_memory::{
mark_usage_capture_memory_omitted, usage_json_heap_estimate, UsageCaptureMemoryBudget,
UsageCaptureRetention,
};
pub use compression::{read_decompressed_usage_json, MAX_DECOMPRESSED_USAGE_JSON_BYTES};
pub use metadata_policy::*;
pub use policy::*;
@@ -673,6 +673,7 @@ mod tests {
fn usage_with_http_capture() -> UpsertUsageRecord {
UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-sensitive-capture".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("key-1".to_string()),
@@ -54,11 +54,11 @@ pub fn extract_provider_reasoning_effort_from_body(value: Option<&Value>) -> Opt
}
fn normalize_provider_reasoning_effort(value: &str) -> Option<String> {
let normalized = value.trim().to_ascii_lowercase();
if normalized.is_empty() || normalized.len() > 64 {
let value = value.trim();
if value.is_empty() || value.len() > 64 {
return None;
}
Some(normalized)
Some(value.to_ascii_lowercase())
}
pub fn extract_provider_service_tier_from_body(value: Option<&Value>) -> Option<String> {
@@ -112,11 +112,11 @@ pub fn extract_provider_actual_service_tier_from_response(value: Option<&Value>)
}
pub fn normalize_provider_service_tier(value: &str) -> Option<String> {
let normalized = value.trim().to_ascii_lowercase();
if normalized.is_empty() || normalized.len() > 64 {
let value = value.trim();
if value.is_empty() || value.len() > 64 {
return None;
}
Some(normalized)
Some(value.to_ascii_lowercase())
}
/// Resolves a provider processing tier exclusively from the final upstream request.
@@ -1980,7 +1980,7 @@ pub trait UsageReadRepository: Send + Sync {
/// Request/response headers and bodies here are capture inputs that the repository persists into
/// the dedicated HTTP audit/body stores. Deprecated mirror columns on `public.usage` remain in the
/// schema for compatibility only and are not the intended long-term destination for new writes.
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
#[derive(Debug, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct UpsertUsageRecord {
pub request_id: String,
pub user_id: Option<String>,
@@ -2055,6 +2055,142 @@ pub struct UpsertUsageRecord {
pub finalized_at_unix_secs: Option<u64>,
pub created_at_unix_ms: Option<u64>,
pub updated_at_unix_secs: u64,
#[doc(hidden)]
#[serde(skip)]
pub capture_retention: super::UsageCaptureRetention,
}
impl Clone for UpsertUsageRecord {
fn clone(&self) -> Self {
let (capture_retention, retain_bodies) = self.capture_retention.clone_for_bodies(|| {
[
&self.request_body,
&self.provider_request_body,
&self.response_body,
&self.client_response_body,
]
.into_iter()
.flatten()
.fold(0usize, |bytes, body| {
bytes
.saturating_add(std::mem::size_of::<Value>())
.saturating_add(super::usage_json_heap_estimate(body))
})
});
let mut cloned = Self {
request_id: self.request_id.clone(),
user_id: self.user_id.clone(),
api_key_id: self.api_key_id.clone(),
username: self.username.clone(),
api_key_name: self.api_key_name.clone(),
provider_name: self.provider_name.clone(),
model: self.model.clone(),
target_model: self.target_model.clone(),
provider_id: self.provider_id.clone(),
provider_endpoint_id: self.provider_endpoint_id.clone(),
provider_api_key_id: self.provider_api_key_id.clone(),
request_type: self.request_type.clone(),
api_format: self.api_format.clone(),
api_family: self.api_family.clone(),
endpoint_kind: self.endpoint_kind.clone(),
endpoint_api_format: self.endpoint_api_format.clone(),
provider_api_family: self.provider_api_family.clone(),
provider_endpoint_kind: self.provider_endpoint_kind.clone(),
has_format_conversion: self.has_format_conversion,
is_stream: self.is_stream,
input_tokens: self.input_tokens,
output_tokens: self.output_tokens,
total_tokens: self.total_tokens,
cache_creation_input_tokens: self.cache_creation_input_tokens,
cache_creation_ephemeral_5m_input_tokens: self.cache_creation_ephemeral_5m_input_tokens,
cache_creation_ephemeral_1h_input_tokens: self.cache_creation_ephemeral_1h_input_tokens,
cache_read_input_tokens: self.cache_read_input_tokens,
cache_creation_cost_usd: self.cache_creation_cost_usd,
cache_read_cost_usd: self.cache_read_cost_usd,
output_price_per_1m: self.output_price_per_1m,
total_cost_usd: self.total_cost_usd,
actual_total_cost_usd: self.actual_total_cost_usd,
status_code: self.status_code,
error_message: self.error_message.clone(),
error_category: self.error_category.clone(),
response_time_ms: self.response_time_ms,
first_byte_time_ms: self.first_byte_time_ms,
status: self.status.clone(),
billing_status: self.billing_status.clone(),
request_headers: self.request_headers.clone(),
request_body: retain_bodies.then(|| self.request_body.clone()).flatten(),
request_body_ref: self.request_body_ref.clone(),
request_body_state: self.request_body_state,
provider_request_headers: self.provider_request_headers.clone(),
provider_request_body: retain_bodies
.then(|| self.provider_request_body.clone())
.flatten(),
provider_request_body_ref: self.provider_request_body_ref.clone(),
provider_request_body_state: self.provider_request_body_state,
response_headers: self.response_headers.clone(),
response_body: retain_bodies.then(|| self.response_body.clone()).flatten(),
response_body_ref: self.response_body_ref.clone(),
response_body_state: self.response_body_state,
client_response_headers: self.client_response_headers.clone(),
client_response_body: retain_bodies
.then(|| self.client_response_body.clone())
.flatten(),
client_response_body_ref: self.client_response_body_ref.clone(),
client_response_body_state: self.client_response_body_state,
candidate_id: self.candidate_id.clone(),
candidate_index: self.candidate_index,
key_name: self.key_name.clone(),
planner_kind: self.planner_kind.clone(),
route_family: self.route_family.clone(),
route_kind: self.route_kind.clone(),
execution_path: self.execution_path.clone(),
local_execution_runtime_miss_reason: self.local_execution_runtime_miss_reason.clone(),
request_metadata: self.request_metadata.clone(),
finalized_at_unix_secs: self.finalized_at_unix_secs,
created_at_unix_ms: self.created_at_unix_ms,
updated_at_unix_secs: self.updated_at_unix_secs,
capture_retention,
};
if !retain_bodies {
for (present, key, state) in [
(
self.request_body.is_some(),
"request",
&mut cloned.request_body_state,
),
(
self.provider_request_body.is_some(),
"provider_request",
&mut cloned.provider_request_body_state,
),
(
self.response_body.is_some(),
"response",
&mut cloned.response_body_state,
),
(
self.client_response_body.is_some(),
"client_response",
&mut cloned.client_response_body_state,
),
] {
if present
&& !matches!(
state,
Some(
UsageBodyCaptureState::None
| UsageBodyCaptureState::Disabled
| UsageBodyCaptureState::Unavailable
)
)
{
*state = Some(UsageBodyCaptureState::Truncated);
super::mark_usage_capture_memory_omitted(&mut cloned.request_metadata, key);
}
}
}
cloned
}
}
impl UpsertUsageRecord {
@@ -2501,15 +2637,59 @@ fn parse_timestamp(value: i64, field_name: &str) -> Result<u64, crate::DataLayer
mod tests {
use super::{
canonical_usage_body_ref_for, extract_provider_actual_service_tier_from_response,
extract_provider_service_tier_from_body, resolve_provider_cache_ttl_minutes,
usage_body_ref, StoredRequestUsageAudit, UpsertUsageRecord, UsageBodyCaptureState,
UsageBodyCaptureStorage, UsageBodyField, UsageProviderPerformanceQuery,
REALTIME_SESSION_METADATA_KEY, USAGE_AVAILABLE_METADATA_KEY,
USAGE_PRICING_AVAILABLE_METADATA_KEY, WEBSOCKET_MODE_METADATA_KEY,
WEBSOCKET_TRANSPORT_METADATA_KEY,
extract_provider_service_tier_from_body, normalize_provider_reasoning_effort,
normalize_provider_service_tier, resolve_provider_cache_ttl_minutes, usage_body_ref,
StoredRequestUsageAudit, UpsertUsageRecord, UsageBodyCaptureState, UsageBodyCaptureStorage,
UsageBodyField, UsageProviderPerformanceQuery, REALTIME_SESSION_METADATA_KEY,
USAGE_AVAILABLE_METADATA_KEY, USAGE_PRICING_AVAILABLE_METADATA_KEY,
WEBSOCKET_MODE_METADATA_KEY, WEBSOCKET_TRANSPORT_METADATA_KEY,
};
use serde_json::{json, Value};
#[test]
fn provider_fact_normalization_preserves_trimmed_byte_limit() {
for normalize in [
normalize_provider_reasoning_effort as fn(&str) -> Option<String>,
normalize_provider_service_tier,
] {
assert_eq!(normalize(" \t\r\n"), None);
assert_eq!(normalize(" HIGH\n"), Some("high".to_string()));
assert_eq!(normalize(&"A".repeat(64)), Some("a".repeat(64)));
assert_eq!(
normalize(&format!(" \t{}\n", "A".repeat(64))),
Some("a".repeat(64))
);
assert_eq!(normalize(&"A".repeat(65)), None);
}
}
#[test]
fn provider_fact_normalization_preserves_non_ascii_case_and_byte_count() {
for normalize in [
normalize_provider_reasoning_effort as fn(&str) -> Option<String>,
normalize_provider_service_tier,
] {
let accepted = format!("{}A", "\u{00c9}".repeat(31));
assert_eq!(
normalize(&accepted),
Some(format!("{}a", "\u{00c9}".repeat(31)))
);
assert_eq!(
normalize(&"\u{00c9}".repeat(32)),
Some("\u{00c9}".repeat(32))
);
assert_eq!(normalize(&format!("{}A", "\u{00c9}".repeat(32))), None);
assert_eq!(normalize("\u{2003}FAST\u{2003}"), Some("fast".to_string()));
}
}
#[test]
fn provider_fact_normalization_rejects_large_input_before_copying() {
let oversized = "A".repeat(4 * 1024 * 1024);
assert_eq!(normalize_provider_reasoning_effort(&oversized), None);
assert_eq!(normalize_provider_service_tier(&oversized), None);
}
fn sample_usage() -> StoredRequestUsageAudit {
StoredRequestUsageAudit::new(
"usage-1".to_string(),
@@ -2700,6 +2880,7 @@ mod tests {
#[test]
fn rejects_invalid_upsert_payload() {
let mut record = UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "".to_string(),
user_id: None,
api_key_id: None,
@@ -16,17 +16,36 @@ impl PostgresBackend {
table_names: &[&str],
) -> Result<DatabaseMaintenanceSummary, DataLayerError> {
let mut summary = DatabaseMaintenanceSummary::default();
if table_names.is_empty() {
return Ok(summary);
}
// VACUUM cannot run inside a transaction. Discard this connection on every
// exit path so its longer session deadlines never leak into request queries.
let mut conn = self.pool().acquire().await.map_postgres_err()?;
conn.close_on_drop();
sqlx::query("SET statement_timeout = '5min'")
.execute(&mut *conn)
.await
.map_postgres_err()?;
sqlx::query("SET lock_timeout = '30s'")
.execute(&mut *conn)
.await
.map_postgres_err()?;
for table_name in table_names {
let table_name = maintenance_identifier(table_name)?;
summary.attempted += 1;
let statement = format!("VACUUM ANALYZE \"{table_name}\"");
if sqlx::raw_sql(&statement)
.execute(self.pool())
match sqlx::query(&statement)
.execute(&mut *conn)
.await
.map_postgres_err()
.is_ok()
{
summary.succeeded += 1;
Ok(_) => summary.succeeded += 1,
Err(error) => tracing::warn!(
table_name,
error = %error,
"PostgreSQL table maintenance failed"
),
}
}
Ok(summary)
@@ -274,6 +274,40 @@ mod tests {
use super::PostgresBackend;
use crate::driver::postgres::{PostgresLeaseRunnerConfig, PostgresPoolConfig};
#[tokio::test]
async fn maintenance_and_aggregation_futures_are_send() {
fn assert_send(_: impl Send) {}
let backend = PostgresBackend::from_config(PostgresPoolConfig {
database_url: "postgres://localhost/aether".to_string(),
min_connections: 0,
..PostgresPoolConfig::default()
})
.unwrap();
let now = chrono::Utc::now();
let daily = crate::StatsDailyAggregationInput {
target_day_utc: now,
aggregated_at: now,
};
let hourly = crate::StatsHourlyAggregationInput {
target_hour_utc: now,
aggregated_at: now,
};
let wallet = crate::WalletDailyUsageAggregationInput {
billing_date: "2026-09-09".to_string(),
billing_timezone: "UTC".to_string(),
window_start_unix_secs: 0,
window_end_unix_secs: 86_400,
aggregated_at_unix_secs: 86_400,
};
// Drop without polling: these are compile-time checks for spawned workers.
assert_send(backend.run_table_maintenance(&["usage"]));
assert_send(backend.aggregate_stats_daily(&daily));
assert_send(backend.aggregate_stats_hourly(&hourly));
assert_send(backend.aggregate_wallet_daily_usage(&wallet));
}
#[tokio::test]
async fn backend_retains_config_and_pool() {
let config = PostgresPoolConfig {
@@ -66,6 +66,12 @@ async fn perform_stats_aggregation_for_day(
) -> Result<StatsDailyAggregationSummary, sqlx::Error> {
let day_end_utc = day_start_utc + chrono::Duration::days(1);
let mut tx = pool.begin().await?;
sqlx::query("SET LOCAL statement_timeout = '5min'")
.execute(&mut *tx)
.await?;
sqlx::query("SET LOCAL lock_timeout = '30s'")
.execute(&mut *tx)
.await?;
let aggregate_row = sqlx::query(SELECT_STATS_DAILY_AGGREGATE_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
@@ -65,6 +65,12 @@ async fn perform_stats_hourly_aggregation_for_hour(
) -> Result<StatsHourlyAggregationSummary, sqlx::Error> {
let hour_end = hour_utc + chrono::Duration::hours(1);
let mut tx = pool.begin().await?;
sqlx::query("SET LOCAL statement_timeout = '5min'")
.execute(&mut *tx)
.await?;
sqlx::query("SET LOCAL lock_timeout = '30s'")
.execute(&mut *tx)
.await?;
let row = sqlx::query(SELECT_STATS_HOURLY_AGGREGATE_SQL)
.bind(hour_utc)
@@ -104,6 +104,14 @@ impl PostgresBackend {
let window_end = unix_secs_to_utc(input.window_end_unix_secs, "window_end")?;
let aggregated_at = unix_secs_to_utc(input.aggregated_at_unix_secs, "aggregated_at")?;
let mut tx = self.pool().begin().await.map_postgres_err()?;
sqlx::query("SET LOCAL statement_timeout = '5min'")
.execute(&mut *tx)
.await
.map_postgres_err()?;
sqlx::query("SET LOCAL lock_timeout = '30s'")
.execute(&mut *tx)
.await
.map_postgres_err()?;
let aggregated_wallets = sqlx::query(UPSERT_WALLET_DAILY_USAGE_LEDGER_SQL)
.bind(window_start)
@@ -57,7 +57,7 @@ pub(super) struct AppliedBackfill {
}
pub async fn run_backfills(pool: &PgPool) -> Result<(), MigrateError> {
let mut conn = pool.acquire().await?;
let mut conn = aether_data_postgres::acquire_postgres_migration_connection(pool).await?;
if BACKFILL_MIGRATOR.locking {
conn.lock().await?;
@@ -45,7 +45,54 @@ fn merge_extra_data(
#[derive(Debug, Default)]
pub struct InMemoryRequestCandidateRepository {
by_id: RwLock<BTreeMap<String, StoredRequestCandidate>>,
rows: RwLock<CandidateRows>,
}
#[derive(Debug, Default)]
struct CandidateRows {
by_id: BTreeMap<String, StoredRequestCandidate>,
by_request: BTreeMap<String, BTreeSet<String>>,
by_created: BTreeSet<(std::cmp::Reverse<u64>, String)>,
}
impl CandidateRows {
fn remove(&mut self, id: &str) -> Option<StoredRequestCandidate> {
let row = self.by_id.remove(id)?;
self.by_created
.remove(&(std::cmp::Reverse(row.created_at_unix_ms), row.id.clone()));
if let Some(ids) = self.by_request.get_mut(&row.request_id) {
ids.remove(id);
if ids.is_empty() {
self.by_request.remove(&row.request_id);
}
}
Some(row)
}
fn insert(&mut self, row: StoredRequestCandidate) -> &StoredRequestCandidate {
// Keep all indexes behind one lock and sanitize every insertion. Reads
// can clone these records without rebuilding their diagnostic JSON.
let row = sanitize_stored_candidate(row);
self.remove(&row.id);
self.by_request
.entry(row.request_id.clone())
.or_default()
.insert(row.id.clone());
self.by_created
.insert((std::cmp::Reverse(row.created_at_unix_ms), row.id.clone()));
self.by_id
.entry(row.id.clone())
.insert_entry(row)
.into_mut()
}
fn for_request(&self, request_id: &str) -> impl Iterator<Item = &StoredRequestCandidate> {
self.by_request
.get(request_id)
.into_iter()
.flatten()
.filter_map(|id| self.by_id.get(id))
}
}
impl InMemoryRequestCandidateRepository {
@@ -53,12 +100,12 @@ impl InMemoryRequestCandidateRepository {
where
I: IntoIterator<Item = StoredRequestCandidate>,
{
let mut by_id = BTreeMap::new();
for item in items.into_iter().map(sanitize_stored_candidate) {
by_id.insert(item.id.clone(), item);
let mut rows = CandidateRows::default();
for item in items {
rows.insert(item);
}
Self {
by_id: RwLock::new(by_id),
rows: RwLock::new(rows),
}
}
}
@@ -70,13 +117,11 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
request_id: &str,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
let mut rows = self
.by_id
.rows
.read()
.expect("request candidate repository lock")
.values()
.filter(|row| row.request_id == request_id)
.for_request(request_id)
.cloned()
.map(sanitize_stored_candidate)
.collect::<Vec<_>>();
rows.sort_by(|left, right| {
left.candidate_index
@@ -95,17 +140,14 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
return Ok(Vec::new());
}
let mut rows = self
.by_id
.read()
.expect("request candidate repository lock")
.values()
let rows = self.rows.read().expect("request candidate repository lock");
Ok(rows
.by_created
.iter()
.take(limit)
.filter_map(|(_, id)| rows.by_id.get(id))
.cloned()
.map(sanitize_stored_candidate)
.collect::<Vec<_>>();
rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms));
rows.truncate(limit);
Ok(rows)
.collect())
}
async fn list_by_provider_id(
@@ -118,19 +160,33 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
}
let mut rows = self
.by_id
.rows
.read()
.expect("request candidate repository lock")
.by_id
.values()
.filter(|row| row.provider_id.as_deref() == Some(provider_id))
.cloned()
.map(sanitize_stored_candidate)
.collect::<Vec<_>>();
rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms));
rows.truncate(limit);
Ok(rows)
}
async fn list_recent_runtime(
&self,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
let rows = self.rows.read().expect("request candidate repository lock");
Ok(rows
.by_created
.iter()
.take(limit)
.filter_map(|(_, id)| rows.by_id.get(id))
.map(StoredRequestCandidate::runtime_snapshot)
.collect())
}
async fn list_finalized_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
@@ -143,9 +199,10 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
let endpoint_ids = endpoint_ids.iter().cloned().collect::<BTreeSet<_>>();
let mut rows = self
.by_id
.rows
.read()
.expect("request candidate repository lock")
.by_id
.values()
.filter(|row| {
row.endpoint_id
@@ -160,7 +217,6 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
)
})
.cloned()
.map(sanitize_stored_candidate)
.collect::<Vec<_>>();
rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms));
rows.truncate(limit);
@@ -179,9 +235,10 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
let endpoint_ids = endpoint_ids.iter().cloned().collect::<BTreeSet<_>>();
let mut counts = BTreeMap::<(String, &'static str), u64>::new();
for row in self
.by_id
.rows
.read()
.expect("request candidate repository lock")
.by_id
.values()
{
let Some(endpoint_id) = row.endpoint_id.as_ref() else {
@@ -243,9 +300,10 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
let mut buckets = BTreeMap::<(String, u32), PublicHealthTimelineBucket>::new();
for row in self
.by_id
.rows
.read()
.expect("request candidate repository lock")
.by_id
.values()
{
let Some(endpoint_id) = row.endpoint_id.as_ref() else {
@@ -319,19 +377,17 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
candidate.sanitize_for_persistence();
candidate.validate()?;
let mut by_id = self
.by_id
let mut rows = self
.rows
.write()
.expect("request candidate repository lock");
let existing = by_id
.values()
let existing = rows
.for_request(&candidate.request_id)
.find(|row| {
row.request_id == candidate.request_id
&& row.candidate_index == candidate.candidate_index
row.candidate_index == candidate.candidate_index
&& row.retry_index == candidate.retry_index
})
.cloned()
.map(sanitize_stored_candidate);
.cloned();
let preserve_existing_lifecycle = existing.as_ref().is_some_and(|row| {
request_candidate_lifecycle_would_regress(row.status, candidate.status)
@@ -449,10 +505,7 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
.or_else(|| existing.as_ref().and_then(|row| row.finished_at_unix_ms))
},
};
let stored = sanitize_stored_candidate(stored);
by_id.insert(stored.id.clone(), stored.clone());
Ok(stored)
Ok(rows.insert(stored).clone())
}
async fn delete_created_before(
@@ -464,11 +517,12 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
return Ok(0);
}
let mut by_id = self
.by_id
let mut rows = self
.rows
.write()
.expect("request candidate repository lock");
let mut ids = by_id
let mut ids = rows
.by_id
.values()
.filter(|row| row.created_at_unix_ms < created_before_unix_secs * 1000)
.map(|row| (row.created_at_unix_ms, row.id.clone()))
@@ -477,7 +531,7 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
let mut deleted = 0usize;
for (_, id) in ids.into_iter().take(limit) {
if by_id.remove(&id).is_some() {
if rows.remove(&id).is_some() {
deleted += 1;
}
}
@@ -546,6 +600,79 @@ mod tests {
assert_eq!(rows[1].request_id, "req-1");
}
#[tokio::test]
async fn request_index_tracks_replaced_ids_and_removes_empty_requests() {
let repository = InMemoryRequestCandidateRepository::seed([
sample_candidate("same-id", "old-request", 100),
sample_candidate("same-id", "new-request", 200),
sample_candidate("other-id", "new-request", 300),
]);
assert!(repository
.list_by_request_id("old-request")
.await
.unwrap()
.is_empty());
assert_eq!(
repository
.list_by_request_id("new-request")
.await
.unwrap()
.len(),
2
);
assert_eq!(repository.delete_created_before(1, 1).await.unwrap(), 1);
let remaining = repository.list_by_request_id("new-request").await.unwrap();
assert_eq!(remaining.len(), 1);
assert_eq!(remaining[0].id, "other-id");
assert_eq!(repository.delete_created_before(1, 1).await.unwrap(), 1);
let rows = repository.rows.read().unwrap();
assert!(rows.by_id.is_empty());
assert!(rows.by_request.is_empty());
assert!(rows.by_created.is_empty());
}
#[tokio::test]
async fn recent_index_preserves_equal_timestamp_order_and_replaced_dates() {
let repository = InMemoryRequestCandidateRepository::seed([
sample_candidate("b", "req-b", 400),
sample_candidate("a", "req-a", 200),
sample_candidate("c", "req-c", 200),
sample_candidate("b", "req-b", 100),
]);
let recent = repository.list_recent(2).await.unwrap();
assert_eq!(
recent.iter().map(|row| row.id.as_str()).collect::<Vec<_>>(),
["a", "c"]
);
assert_eq!(repository.list_recent(0).await.unwrap().len(), 0);
assert_eq!(repository.rows.read().unwrap().by_created.len(), 3);
}
#[tokio::test]
async fn runtime_reads_keep_metadata_without_diagnostic_payloads() {
let mut candidate = sample_candidate("candidate", "request", 100);
candidate.extra_data = Some(json!({"upstream_response": {"body": "x".repeat(32_768)}}));
candidate.error_message = Some("diagnostic detail".into());
candidate.required_capabilities = Some(json!({"vision": true}));
candidate.concurrent_requests = Some(17);
let repository = InMemoryRequestCandidateRepository::seed([candidate]);
let full = repository.list_recent(1).await.unwrap();
let runtime = repository.list_recent_runtime(1).await.unwrap();
assert_eq!(
runtime,
full.iter()
.map(StoredRequestCandidate::runtime_snapshot)
.collect::<Vec<_>>()
);
assert_eq!(runtime[0].concurrent_requests, Some(17));
assert!(runtime[0].extra_data.is_none());
assert!(runtime[0].error_message.is_none());
assert!(runtime[0].required_capabilities.is_none());
assert!(full[0].extra_data.is_some());
assert_eq!(repository.list_recent(1).await.unwrap(), full);
assert!(repository.list_recent_runtime(0).await.unwrap().is_empty());
}
#[tokio::test]
async fn lists_recent_request_candidates_in_descending_created_order() {
let repository = InMemoryRequestCandidateRepository::seed(vec![
@@ -602,10 +729,11 @@ mod tests {
{
let stored = repository
.by_id
.rows
.read()
.expect("request candidate repository lock");
let candidate = stored
.by_id
.get("cand-raw")
.expect("seeded candidate should exist");
assert_eq!(
@@ -628,10 +756,10 @@ mod tests {
bypassed_candidate.id = "cand-bypassed".to_string();
bypassed_candidate.request_id = "req-bypassed".to_string();
repository
.by_id
.rows
.write()
.expect("request candidate repository lock")
.insert(bypassed_candidate.id.clone(), bypassed_candidate);
.insert(bypassed_candidate);
let rows = repository
.list_recent(10)
@@ -66,6 +66,7 @@ fn sample_usage(request_id: &str, created_at_unix_ms: i64) -> StoredRequestUsage
fn sample_upsert_usage_record(request_id: &str) -> UpsertUsageRecord {
UpsertUsageRecord {
capture_retention: Default::default(),
request_id: request_id.to_string(),
user_id: None,
api_key_id: None,
@@ -535,6 +536,7 @@ async fn stale_pending_update_does_not_regress_finalized_usage() {
let repository = InMemoryUsageReadRepository::default();
repository
.upsert(UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-finalized-1".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("api-key-1".to_string()),
@@ -608,6 +610,7 @@ async fn stale_pending_update_does_not_regress_finalized_usage() {
repository
.upsert(UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-finalized-1".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("api-key-1".to_string()),
@@ -695,6 +698,7 @@ async fn upsert_allows_completed_recovery_after_void_failure() {
let repository = InMemoryUsageReadRepository::default();
repository
.upsert(UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-recover-1".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("api-key-1".to_string()),
@@ -768,6 +772,7 @@ async fn upsert_allows_completed_recovery_after_void_failure() {
repository
.upsert(UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-recover-1".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("api-key-1".to_string()),
@@ -1009,6 +1014,7 @@ async fn stale_pending_update_does_not_regress_streaming_usage() {
let repository = InMemoryUsageReadRepository::default();
repository
.upsert(UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-streaming-1".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("api-key-1".to_string()),
@@ -1084,6 +1090,7 @@ async fn stale_pending_update_does_not_regress_streaming_usage() {
repository
.upsert(UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-streaming-1".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("api-key-1".to_string()),
@@ -1339,6 +1346,7 @@ async fn upsert_writes_usage_record() {
let repository = InMemoryUsageReadRepository::default();
let stored = repository
.upsert(UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-upsert-1".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("key-1".to_string()),
@@ -1430,6 +1438,7 @@ async fn upsert_defaults_created_at_to_second_timestamp() {
let repository = InMemoryUsageReadRepository::default();
let stored = repository
.upsert(UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-upsert-ms-default".to_string(),
user_id: None,
api_key_id: None,
@@ -1509,6 +1518,7 @@ async fn upsert_does_not_backfill_legacy_output_price_from_request_metadata() {
let repository = InMemoryUsageReadRepository::default();
let stored = repository
.upsert(UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-upsert-price-metadata".to_string(),
user_id: None,
api_key_id: None,
@@ -1591,6 +1601,7 @@ async fn upsert_does_not_backfill_typed_body_refs_from_request_metadata() {
let repository = InMemoryUsageReadRepository::default();
let stored = repository
.upsert(UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-upsert-body-ref-metadata".to_string(),
user_id: None,
api_key_id: None,
@@ -1673,6 +1684,7 @@ async fn upsert_keeps_typed_routing_fields_out_of_request_metadata() {
let repository = InMemoryUsageReadRepository::default();
let stored = repository
.upsert(UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-upsert-routing-metadata".to_string(),
user_id: None,
api_key_id: None,
@@ -1771,6 +1783,7 @@ async fn upsert_does_not_persist_legacy_display_columns_for_new_rows() {
let repository = InMemoryUsageReadRepository::default();
let stored = repository
.upsert(UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-upsert-display-columns".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("key-1".to_string()),
@@ -1921,6 +1934,7 @@ async fn upsert_preserves_existing_legacy_display_columns_when_new_write_omits_t
}]);
let stored = repository
.upsert(UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-existing-display-columns".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("key-1".to_string()),
@@ -57,6 +57,7 @@ mod tests {
#[test]
fn strip_deprecated_usage_display_fields_clears_legacy_display_columns() {
let usage = strip_deprecated_usage_display_fields(UpsertUsageRecord {
capture_retention: Default::default(),
request_id: "req-1".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("key-1".to_string()),
+7 -2
View File
@@ -10,13 +10,18 @@ description = "HTTP frontdoor middleware and request lifecycle primitives for Ae
aether-ai-formats.workspace = true
axum.workspace = true
bytes.workspace = true
futures-util.workspace = true
http.workspace = true
tokio.workspace = true
tokio = { workspace = true, features = ["io-util"] }
tokio-util = { workspace = true, features = ["rt"] }
tracing.workspace = true
uuid.workspace = true
[dev-dependencies]
futures-util.workspace = true
http-body-util = "0.1"
hyper = { version = "1", features = ["client", "server", "http1", "http2"] }
hyper-util = { version = "0.1", features = ["tokio"] }
serde_json.workspace = true
tokio = { workspace = true, features = ["test-util"] }
tower = { version = "0.5", features = ["util"] }
tracing-subscriber.workspace = true
+303 -75
View File
@@ -1,14 +1,13 @@
//! Bounded request-body buffering for frontdoor adapters.
//!
//! The policy reserves weighted memory before reading a body and holds the
//! reservation through the caller's normalization callback. This keeps body
//! buffering independent from gateway business routing while preventing a
//! burst of compressed requests from bypassing the memory budget.
//! The policy grows weighted reservations as bytes arrive and holds them through
//! normalization. Growth never waits while retaining a partial buffer, so
//! concurrent uploads cannot deadlock while competing for the remaining budget.
use axum::body::{to_bytes, Body};
use axum::body::Body;
use bytes::Bytes;
use futures_util::StreamExt;
use http::{header, HeaderMap, StatusCode};
use std::error::Error as StdError;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
@@ -127,7 +126,12 @@ impl BodyBufferPolicy {
}
pub fn reservation_bytes(&self, headers: &HeaderMap) -> usize {
reservation_bytes(headers, self.max_bytes, self.budget_bytes)
reservation_bytes(
headers,
self.max_bytes,
self.budget_bytes,
self.permit_bytes,
)
}
pub fn reservation_permits(&self, reservation_bytes: usize) -> u32 {
@@ -168,67 +172,134 @@ impl BodyBufferPolicy {
};
Ok(BodyBufferReservation {
permit,
memory: BodyBufferBudget {
permit,
budget: Arc::clone(&self.budget),
budget_bytes: self.budget_bytes,
permit_bytes: self.permit_bytes,
requested_bytes,
},
max_bytes: effective_max_bytes,
read_timeout: self.read_timeout,
requested_bytes,
})
}
}
#[derive(Debug)]
pub struct BodyBufferReservation {
permit: OwnedSemaphorePermit,
memory: BodyBufferBudget,
max_bytes: u64,
read_timeout: Option<Duration>,
}
#[derive(Debug)]
pub struct BodyBufferBudget {
permit: OwnedSemaphorePermit,
budget: Arc<Semaphore>,
budget_bytes: usize,
permit_bytes: usize,
requested_bytes: usize,
}
impl BodyBufferBudget {
/// Reserve a new high-water mark before retaining or decoding more bytes.
/// Never queue for growth while another partial request may hold the rest.
pub fn try_reserve_bytes(&mut self, requested_bytes: usize) -> Result<(), BodyBufferError> {
if requested_bytes > self.budget_bytes {
return Err(self.overloaded(requested_bytes));
}
let permits = reservation_permits(requested_bytes, self.permit_bytes) as usize;
let additional = permits.saturating_sub(self.permit.num_permits());
if additional > 0 {
let permit = Arc::clone(&self.budget)
.try_acquire_many_owned(additional as u32)
.map_err(|_| self.overloaded(requested_bytes))?;
self.permit.merge(permit);
}
self.requested_bytes = self.requested_bytes.max(requested_bytes);
Ok(())
}
fn overloaded(&self, requested_bytes: usize) -> BodyBufferError {
BodyBufferError::Overloaded {
requested_bytes,
budget_bytes: self.budget_bytes,
timeout_ms: 0,
}
}
}
impl BodyBufferReservation {
pub fn requested_bytes(&self) -> usize {
self.requested_bytes
self.memory.requested_bytes
}
pub async fn collect(self, body: Body) -> Result<BufferedBody, BodyBufferError> {
let Self {
permit,
mut memory,
max_bytes,
read_timeout,
requested_bytes,
} = self;
let started_at = Instant::now();
let body_limit = usize::try_from(max_bytes).unwrap_or(usize::MAX);
let collected = match read_timeout {
Some(read_timeout) => {
match tokio::time::timeout(read_timeout, to_bytes(body, body_limit)).await {
Ok(result) => result,
Err(_) => {
return Err(BodyBufferError::Timeout {
timeout_ms: duration_millis(read_timeout),
});
let collect = async {
let mut stream = body.into_data_stream();
let mut bytes = Vec::new();
while let Some(chunk) = stream.next().await {
let chunk = chunk.map_err(|error| BodyBufferError::ReadFailed {
message: error.to_string(),
})?;
let length =
bytes
.len()
.checked_add(chunk.len())
.ok_or(BodyBufferError::TooLarge {
limit_bytes: max_bytes,
})?;
if length as u64 > max_bytes {
return Err(BodyBufferError::TooLarge {
limit_bytes: max_bytes,
});
}
if length > bytes.capacity() {
let mut capacity = if length <= DEFAULT_BODY_BUFFER_PERMIT_BYTES {
length
} else {
bytes
.capacity()
.saturating_mul(2)
.max(length)
.min(usize::try_from(max_bytes).unwrap_or(usize::MAX))
};
if memory.try_reserve_bytes(capacity).is_err() {
memory.try_reserve_bytes(length)?;
capacity = length;
}
// Account for geometric growth, falling back to the bytes needed under load.
bytes
.try_reserve_exact(capacity - bytes.len())
.map_err(|error| BodyBufferError::ReadFailed {
message: error.to_string(),
})?;
if bytes.capacity() > capacity {
memory.try_reserve_bytes(bytes.capacity())?;
}
}
bytes.extend_from_slice(&chunk);
}
None => to_bytes(body, body_limit).await,
};
let bytes = match collected {
Ok(bytes) => bytes,
Err(error) if collection_exceeded_limit(&error) => {
return Err(BodyBufferError::TooLarge {
limit_bytes: max_bytes,
});
}
Err(error) => {
return Err(BodyBufferError::ReadFailed {
message: error.to_string(),
});
}
Ok(Bytes::from(bytes))
};
let bytes = match read_timeout {
Some(read_timeout) => tokio::time::timeout(read_timeout, collect)
.await
.map_err(|_| BodyBufferError::Timeout {
timeout_ms: duration_millis(read_timeout),
}),
None => Ok(collect.await),
}??;
Ok(BufferedBody {
bytes,
permit: Some(permit),
requested_bytes,
memory,
elapsed: started_at.elapsed(),
})
}
@@ -237,8 +308,7 @@ impl BodyBufferReservation {
#[derive(Debug)]
pub struct BufferedBody {
bytes: Bytes,
permit: Option<OwnedSemaphorePermit>,
requested_bytes: usize,
memory: BodyBufferBudget,
elapsed: Duration,
}
@@ -248,19 +318,28 @@ impl BufferedBody {
}
pub fn requested_bytes(&self) -> usize {
self.requested_bytes
self.memory.requested_bytes
}
pub fn elapsed(&self) -> Duration {
self.elapsed
}
/// Apply normalization while retaining the memory permit until the
/// callback completes.
/// Retain the permit for a callback that does not grow the buffered payload.
pub fn try_map<T, E>(self, map: impl FnOnce(Bytes) -> Result<T, E>) -> Result<T, E> {
let Self { bytes, permit, .. } = self;
let result = map(bytes);
drop(permit);
self.try_map_with_budget(|bytes, _| map(bytes))
}
/// The callback must account for decoded buffers before allocating them.
pub fn try_map_with_budget<T, E>(
self,
map: impl FnOnce(Bytes, &mut BodyBufferBudget) -> Result<T, E>,
) -> Result<T, E> {
let Self {
bytes, mut memory, ..
} = self;
let result = map(bytes, &mut memory);
drop(memory);
result
}
}
@@ -389,23 +468,15 @@ fn invalid_body_headers(message: &str) -> BodyBufferError {
}
}
fn reservation_bytes(headers: &HeaderMap, max_bytes: u64, budget_bytes: usize) -> usize {
fn reservation_bytes(
headers: &HeaderMap,
max_bytes: u64,
budget_bytes: usize,
permit_bytes: usize,
) -> usize {
let reservation_ceiling = usize::try_from(max_bytes)
.unwrap_or(usize::MAX)
.min(budget_bytes);
let encoded = headers
.get_all(header::CONTENT_ENCODING)
.iter()
.any(|value| {
value.to_str().map_or(true, |value| {
value.split(',').map(str::trim).any(|encoding| {
!encoding.is_empty() && !encoding.eq_ignore_ascii_case("identity")
})
})
});
if encoded {
return reservation_ceiling;
}
declared_content_length(headers)
.ok()
.flatten()
@@ -414,7 +485,7 @@ fn reservation_bytes(headers: &HeaderMap, max_bytes: u64, budget_bytes: usize) -
.unwrap_or(usize::MAX)
.min(reservation_ceiling)
})
.unwrap_or(reservation_ceiling)
.unwrap_or_else(|| permit_bytes.min(reservation_ceiling))
}
fn reservation_permits(reservation_bytes: usize, permit_bytes: usize) -> u32 {
@@ -429,17 +500,6 @@ fn duration_millis(duration: Duration) -> u64 {
u64::try_from(duration.as_millis()).unwrap_or(u64::MAX)
}
fn collection_exceeded_limit(error: &(dyn StdError + 'static)) -> bool {
let mut current = Some(error);
while let Some(error) = current {
if error.to_string().contains("length limit exceeded") {
return true;
}
current = error.source();
}
false
}
#[cfg(test)]
mod tests {
use super::{BodyBufferError, BodyBufferPolicy, DEFAULT_BODY_BUFFER_PERMIT_BYTES};
@@ -579,9 +639,9 @@ mod tests {
let reservation = policy
.reserve(&headers)
.await
.expect("encoded unlimited body should reserve the available budget");
assert_eq!(reservation.requested_bytes(), 4);
assert_eq!(budget.available_permits(), 0);
.expect("encoded unlimited body should reserve its initial chunk");
assert_eq!(reservation.requested_bytes(), 1);
assert_eq!(budget.available_permits(), 3);
let error = reservation
.collect(Body::from(Bytes::from_static(b"01234")))
@@ -707,4 +767,172 @@ mod tests {
.expect_err("exhausted budget should fail closed");
assert!(matches!(error, BodyBufferError::Overloaded { .. }));
}
#[tokio::test]
async fn small_compressed_and_unknown_length_requests_share_the_budget() {
let budget_bytes = 256 * 1024 * 1024;
let budget = Arc::new(Semaphore::new(
budget_bytes / DEFAULT_BODY_BUFFER_PERMIT_BYTES,
));
let policy = policy(
budget_bytes as u64,
Duration::from_secs(1),
Arc::clone(&budget),
);
let mut headers = HeaderMap::new();
headers.insert(header::CONTENT_ENCODING, HeaderValue::from_static("gzip"));
headers.insert(header::CONTENT_LENGTH, HeaderValue::from_static("1024"));
let compressed = policy.reserve(&headers).await.unwrap();
let unknown = policy.reserve(&HeaderMap::new()).await.unwrap();
assert_eq!(compressed.requested_bytes(), 1024);
assert_eq!(unknown.requested_bytes(), DEFAULT_BODY_BUFFER_PERMIT_BYTES);
assert_eq!(budget.available_permits(), 4094);
drop((compressed, unknown));
assert_eq!(budget.available_permits(), 4096);
}
#[tokio::test]
async fn unknown_length_body_grows_its_reservation_and_holds_it_until_normalized() {
let budget = Arc::new(Semaphore::new(8));
let policy = BodyBufferPolicy::with_permit_bytes(
8,
Duration::from_secs(1),
Duration::from_secs(1),
8,
1,
Arc::clone(&budget),
);
let reservation = policy.reserve(&HeaderMap::new()).await.unwrap();
assert_eq!(budget.available_permits(), 7);
let body = Body::from_stream(stream::iter([
Ok::<_, std::io::Error>(Bytes::from_static(b"ab")),
Ok(Bytes::from_static(b"cd")),
]));
let buffered = reservation.collect(body).await.unwrap();
assert_eq!(buffered.bytes().as_ref(), b"abcd");
assert_eq!(budget.available_permits(), 4);
buffered
.try_map_with_budget(|_, memory| {
memory.try_reserve_bytes(7)?;
assert_eq!(budget.available_permits(), 1);
Ok::<_, BodyBufferError>(())
})
.unwrap();
assert_eq!(budget.available_permits(), 8);
}
#[tokio::test]
async fn partial_upload_growth_rejects_without_waiting_and_releases_its_budget() {
let budget = Arc::new(Semaphore::new(2));
let policy = BodyBufferPolicy::with_permit_bytes(
2,
Duration::from_secs(60),
Duration::from_secs(60),
2,
1,
Arc::clone(&budget),
);
let first = policy.reserve(&HeaderMap::new()).await.unwrap();
let second = policy.reserve(&HeaderMap::new()).await.unwrap();
let error =
tokio::time::timeout(Duration::from_millis(100), first.collect(Body::from("ab")))
.await
.expect("growth must not wait while holding a partial reservation")
.unwrap_err();
assert_eq!(
error,
BodyBufferError::Overloaded {
requested_bytes: 2,
budget_bytes: 2,
timeout_ms: 0,
}
);
assert_eq!(budget.available_permits(), 1);
let buffered = second.collect(Body::from("ab")).await.unwrap();
assert_eq!(budget.available_permits(), 0);
drop(buffered);
assert_eq!(budget.available_permits(), 2);
}
#[tokio::test]
async fn normalization_budget_failure_releases_all_upload_permits() {
let budget = Arc::new(Semaphore::new(4));
let policy = BodyBufferPolicy::with_permit_bytes(
4,
Duration::from_secs(1),
Duration::from_secs(1),
4,
1,
Arc::clone(&budget),
);
let buffered = policy
.reserve(&HeaderMap::new())
.await
.unwrap()
.collect(Body::from("ab"))
.await
.unwrap();
let result = buffered.try_map_with_budget(|_, memory| memory.try_reserve_bytes(5));
assert!(matches!(
result,
Err(BodyBufferError::Overloaded {
requested_bytes: 5,
..
})
));
assert_eq!(budget.available_permits(), 4);
}
#[tokio::test]
async fn upload_growth_uses_available_budget_without_requiring_geometric_headroom() {
let budget = Arc::new(Semaphore::new(128));
let held = Arc::clone(&budget).acquire_many_owned(32).await.unwrap();
let policy = BodyBufferPolicy::with_permit_bytes(
128 * 1024,
Duration::from_secs(1),
Duration::from_secs(1),
128 * 1024,
1024,
Arc::clone(&budget),
);
let body = Body::from_stream(stream::iter([
Ok::<_, std::io::Error>(Bytes::from(vec![b'a'; 70_000])),
Ok(Bytes::from(vec![b'b'; 20_000])),
]));
let buffered = policy
.reserve(&HeaderMap::new())
.await
.unwrap()
.collect(body)
.await
.unwrap();
assert_eq!(buffered.bytes().len(), 90_000);
assert_eq!(buffered.requested_bytes(), 90_000);
drop((buffered, held));
assert_eq!(budget.available_permits(), 128);
}
#[tokio::test]
async fn cancellation_releases_partial_upload_budget() {
let budget = Arc::new(Semaphore::new(4));
let policy = BodyBufferPolicy::with_permit_bytes(
4,
Duration::from_secs(60),
Duration::from_secs(1),
4,
1,
Arc::clone(&budget),
);
let reservation = policy.reserve(&HeaderMap::new()).await.unwrap();
let body = Body::from_stream(
stream::once(async { Ok::<_, std::io::Error>(Bytes::from_static(b"ab")) })
.chain(stream::pending()),
);
assert!(
tokio::time::timeout(Duration::from_millis(10), reservation.collect(body))
.await
.is_err()
);
assert_eq!(budget.available_permits(), 4);
}
}
@@ -0,0 +1,229 @@
use std::future::Future;
use std::io;
use std::net::SocketAddr;
use std::pin::Pin;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Duration;
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
use tokio_util::sync::{CancellationToken, WaitForCancellationFutureOwned};
const MAX_HTTP_CONNECTIONS: usize = 65_536;
const FD_RESERVE: usize = 256;
/// Selects a per-server incoming TCP limit. The FD allowance leaves room for
/// upstream sockets and process infrastructure; it is not a complete FD budget.
pub fn http_connection_limit(
configured: Option<usize>,
request_limit: usize,
websocket_limit: usize,
fd_soft_limit: Option<usize>,
) -> usize {
let configured = configured
.filter(|limit| *limit > 0)
.unwrap_or_else(|| request_limit.saturating_add(websocket_limit))
.clamp(1, MAX_HTTP_CONNECTIONS);
let fd_allowance = fd_soft_limit
.map(|limit| (limit.saturating_sub(FD_RESERVE) / 2).max(1))
.unwrap_or(MAX_HTTP_CONNECTIONS);
configured.min(fd_allowance)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct HttpConnectionBudgetSnapshot {
pub limit: usize,
pub in_flight: usize,
pub high_watermark: usize,
pub rejected_total: u64,
pub accept_errors_total: u64,
}
/// Share one budget across listeners. Admission belongs to the underlying IO,
/// so HTTP/1 upgrades keep their permit and HTTP/2 streams share one permit.
#[derive(Debug)]
pub struct HttpConnectionBudget {
limit: usize,
permits: Arc<Semaphore>,
in_flight: AtomicUsize,
high_watermark: AtomicUsize,
rejected_total: AtomicU64,
accept_errors_total: AtomicU64,
shutdown: CancellationToken,
}
impl HttpConnectionBudget {
pub fn new(limit: usize) -> Self {
let limit = limit.clamp(1, MAX_HTTP_CONNECTIONS.min(Semaphore::MAX_PERMITS));
Self {
limit,
permits: Arc::new(Semaphore::new(limit)),
in_flight: AtomicUsize::new(0),
high_watermark: AtomicUsize::new(0),
rejected_total: AtomicU64::new(0),
accept_errors_total: AtomicU64::new(0),
shutdown: CancellationToken::new(),
}
}
/// Admit after accepting. Waiting for a permit before accept can let idle
/// reuseport listeners monopolize permits needed by a busy listener.
pub fn try_admit<T>(self: &Arc<Self>, io: T) -> Result<AdmittedConnection<T>, ()> {
let permit = Arc::clone(&self.permits).try_acquire_owned().map_err(|_| {
self.rejected_total.fetch_add(1, Ordering::Relaxed);
})?;
let in_flight = self.in_flight.fetch_add(1, Ordering::Relaxed) + 1;
self.high_watermark.fetch_max(in_flight, Ordering::Relaxed);
Ok(AdmittedConnection {
io,
read_shutdown: Box::pin(self.shutdown.clone().cancelled_owned()),
write_shutdown: Box::pin(self.shutdown.clone().cancelled_owned()),
_permit: ConnectionPermit {
budget: Arc::clone(self),
_permit: permit,
},
})
}
pub fn snapshot(&self) -> HttpConnectionBudgetSnapshot {
HttpConnectionBudgetSnapshot {
limit: self.limit,
in_flight: self.in_flight.load(Ordering::Relaxed),
high_watermark: self.high_watermark.load(Ordering::Relaxed),
rejected_total: self.rejected_total.load(Ordering::Relaxed),
accept_errors_total: self.accept_errors_total.load(Ordering::Relaxed),
}
}
/// End the drain deadline for all sockets, including upgraded connections.
pub fn force_close(&self) {
self.permits.close();
self.shutdown.cancel();
}
pub async fn wait_for_forced_close(&self) {
self.shutdown.cancelled().await;
}
pub async fn accept(&self, listener: &TcpListener) -> (TcpStream, SocketAddr) {
self.accept_with(|| listener.accept()).await
}
async fn accept_with<T, F, A>(&self, mut accept: A) -> T
where
A: FnMut() -> F,
F: Future<Output = io::Result<T>>,
{
loop {
match accept().await {
Ok(connection) => return connection,
Err(error) => {
self.accept_errors_total.fetch_add(1, Ordering::Relaxed);
// Match Axum's listener behavior: failed peers can be retried
// immediately; resource failures such as EMFILE need backoff.
if matches!(
error.kind(),
io::ErrorKind::ConnectionRefused
| io::ErrorKind::ConnectionAborted
| io::ErrorKind::ConnectionReset
) {
continue;
}
tracing::error!(
event_name = "http_connection_accept_failed",
error = %error,
"HTTP listener accept failed; retrying after one second"
);
tokio::time::sleep(Duration::from_secs(1)).await;
}
}
}
}
}
#[derive(Debug)]
struct ConnectionPermit {
budget: Arc<HttpConnectionBudget>,
_permit: OwnedSemaphorePermit,
}
impl Drop for ConnectionPermit {
fn drop(&mut self) {
self.budget.in_flight.fetch_sub(1, Ordering::Relaxed);
}
}
#[derive(Debug)]
pub struct AdmittedConnection<T> {
// Close the socket before returning its permit, including upgrade teardown.
io: T,
read_shutdown: Pin<Box<WaitForCancellationFutureOwned>>,
write_shutdown: Pin<Box<WaitForCancellationFutureOwned>>,
_permit: ConnectionPermit,
}
impl<T: AsyncRead + Unpin> AsyncRead for AdmittedConnection<T> {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
if self.read_shutdown.as_mut().poll(cx).is_ready() {
return Poll::Ready(Err(shutdown_error()));
}
Pin::new(&mut self.io).poll_read(cx, buf)
}
}
impl<T: AsyncWrite + Unpin> AsyncWrite for AdmittedConnection<T> {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
if self.write_shutdown.as_mut().poll(cx).is_ready() {
return Poll::Ready(Err(shutdown_error()));
}
Pin::new(&mut self.io).poll_write(cx, buf)
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
if self.write_shutdown.as_mut().poll(cx).is_ready() {
return Poll::Ready(Err(shutdown_error()));
}
Pin::new(&mut self.io).poll_flush(cx)
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.io).poll_shutdown(cx)
}
fn is_write_vectored(&self) -> bool {
self.io.is_write_vectored()
}
fn poll_write_vectored(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
bufs: &[io::IoSlice<'_>],
) -> Poll<io::Result<usize>> {
if self.write_shutdown.as_mut().poll(cx).is_ready() {
return Poll::Ready(Err(shutdown_error()));
}
Pin::new(&mut self.io).poll_write_vectored(cx, bufs)
}
}
fn shutdown_error() -> io::Error {
io::Error::new(
io::ErrorKind::ConnectionAborted,
"gateway shutdown deadline reached",
)
}
#[cfg(test)]
#[path = "connection_tests.rs"]
mod tests;
@@ -0,0 +1,411 @@
use std::collections::VecDeque;
use std::convert::Infallible;
use std::sync::atomic::AtomicBool;
use bytes::Bytes;
use http::{Request, Response, StatusCode};
use http_body_util::{BodyExt, Empty};
use hyper::body::Incoming;
use hyper::service::service_fn;
use hyper_util::rt::{TokioExecutor, TokioIo, TokioTimer};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use super::*;
async fn within<T>(future: impl Future<Output = T>) -> T {
tokio::time::timeout(Duration::from_secs(3), future)
.await
.expect("connection test exceeded its deadline")
}
async fn tcp_pair() -> (TcpStream, TcpStream) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let client = TcpStream::connect(listener.local_addr().unwrap())
.await
.unwrap();
let (server, _) = listener.accept().await.unwrap();
(client, server)
}
#[tokio::test]
async fn forced_shutdown_wakes_independent_read_and_write_waiters() {
let budget = Arc::new(HttpConnectionBudget::new(2));
let (io, _peer) = tokio::io::duplex(1);
let io = budget.try_admit(io).unwrap();
let (mut reader, mut writer) = tokio::io::split(io);
writer.write_all(b"a").await.unwrap();
let reading = tokio::spawn(async move { reader.read_u8().await });
let writing = tokio::spawn(async move { writer.write_all(b"b").await });
tokio::task::yield_now().await;
budget.force_close();
assert_eq!(
within(reading).await.unwrap().unwrap_err().kind(),
io::ErrorKind::ConnectionAborted
);
assert_eq!(
within(writing).await.unwrap().unwrap_err().kind(),
io::ErrorKind::ConnectionAborted
);
assert_eq!(budget.snapshot().in_flight, 0);
assert!(budget.try_admit(tokio::io::empty()).is_err());
}
#[tokio::test]
async fn forced_shutdown_before_first_poll_closes_read_write_and_flush() {
let budget = Arc::new(HttpConnectionBudget::new(1));
let (io, _peer) = tokio::io::duplex(16);
let mut io = budget.try_admit(io).unwrap();
budget.force_close();
assert_eq!(
io.read_u8().await.unwrap_err().kind(),
io::ErrorKind::ConnectionAborted
);
assert_eq!(
io.write_all(b"a").await.unwrap_err().kind(),
io::ErrorKind::ConnectionAborted
);
assert_eq!(
io.flush().await.unwrap_err().kind(),
io::ErrorKind::ConnectionAborted
);
let buffers = [io::IoSlice::new(b"a")];
assert_eq!(
io.write_vectored(&buffers).await.unwrap_err().kind(),
io::ErrorKind::ConnectionAborted
);
io.shutdown().await.unwrap();
drop(io);
assert_eq!(budget.snapshot().in_flight, 0);
}
#[test]
fn http_connection_limits_apply_to_auto_explicit_zero_and_fd_bounds() {
for (configured, requests, websockets, fd_limit, expected) in [
(None, 0, 0, None, 1),
(None, 3, 5, None, 8),
(Some(0), 3, 5, None, 8),
(Some(9), 3, 5, None, 9),
(Some(usize::MAX), 0, 0, None, MAX_HTTP_CONNECTIONS),
(None, usize::MAX, usize::MAX, None, MAX_HTTP_CONNECTIONS),
(None, 512, 512, Some(1_024), 384),
(Some(500), 3, 5, Some(1_024), 384),
(Some(32), 3, 5, Some(1_024), 32),
(Some(500), 3, 5, Some(260), 2),
(Some(500), 3, 5, Some(258), 1),
(Some(500), 3, 5, Some(0), 1),
] {
assert_eq!(
http_connection_limit(configured, requests, websockets, fd_limit),
expected,
);
}
assert_eq!(HttpConnectionBudget::new(0).snapshot().limit, 1);
assert_eq!(
HttpConnectionBudget::new(usize::MAX).snapshot().limit,
MAX_HTTP_CONNECTIONS,
);
}
#[test]
fn http_connection_budget_drops_io_before_returning_capacity() {
struct DropProbe {
budget: Arc<HttpConnectionBudget>,
observed: Arc<AtomicBool>,
}
impl Drop for DropProbe {
fn drop(&mut self) {
self.observed
.store(self.budget.snapshot().in_flight == 1, Ordering::Relaxed);
}
}
let budget = Arc::new(HttpConnectionBudget::new(1));
let admitted_dropped = Arc::new(AtomicBool::new(false));
let rejected_dropped = Arc::new(AtomicBool::new(false));
let admitted = budget
.try_admit(DropProbe {
budget: Arc::clone(&budget),
observed: Arc::clone(&admitted_dropped),
})
.unwrap();
assert!(budget
.try_admit(DropProbe {
budget: Arc::clone(&budget),
observed: Arc::clone(&rejected_dropped),
})
.is_err());
assert!(rejected_dropped.load(Ordering::Relaxed));
assert_eq!(budget.snapshot().in_flight, 1);
drop(admitted);
assert!(admitted_dropped.load(Ordering::Relaxed));
assert_eq!(budget.snapshot().in_flight, 0);
assert_eq!(budget.snapshot().high_watermark, 1);
assert_eq!(budget.snapshot().rejected_total, 1);
drop(budget.try_admit(()).unwrap());
assert_eq!(budget.snapshot().in_flight, 0);
}
#[tokio::test]
async fn http_connection_budget_forwards_read_write_vectored_and_half_close() {
let budget = Arc::new(HttpConnectionBudget::new(1));
let (mut peer, io) = tokio::io::duplex(32);
let vectored = io.is_write_vectored();
let mut admitted = budget.try_admit(io).unwrap();
assert_eq!(admitted.is_write_vectored(), vectored);
let written = admitted
.write_vectored(&[io::IoSlice::new(b"ab"), io::IoSlice::new(b"cd")])
.await
.unwrap();
assert!((1..=4).contains(&written));
admitted.write_all(&b"abcd"[written..]).await.unwrap();
admitted.flush().await.unwrap();
let mut message = [0; 4];
peer.read_exact(&mut message).await.unwrap();
assert_eq!(&message, b"abcd");
admitted.shutdown().await.unwrap();
assert_eq!(peer.read(&mut [0; 1]).await.unwrap(), 0);
assert_eq!(budget.snapshot().in_flight, 1);
peer.write_all(b"reply").await.unwrap();
let mut reply = [0; 5];
admitted.read_exact(&mut reply).await.unwrap();
assert_eq!(&reply, b"reply");
drop(admitted);
assert_eq!(budget.snapshot().in_flight, 0);
}
#[tokio::test]
async fn http_connection_budget_two_tcp_listeners_share_capacity_and_recover() {
let first_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let second_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let budget = Arc::new(HttpConnectionBudget::new(1));
let first_client = TcpStream::connect(first_listener.local_addr().unwrap())
.await
.unwrap();
let (first_io, _) = within(budget.accept(&first_listener)).await;
let first = budget.try_admit(first_io).unwrap();
let mut rejected_client = TcpStream::connect(second_listener.local_addr().unwrap())
.await
.unwrap();
let (second_io, _) = within(budget.accept(&second_listener)).await;
assert!(Arc::clone(&budget).try_admit(second_io).is_err());
let rejected_read = within(rejected_client.read(&mut [0; 1])).await;
assert!(
matches!(rejected_read, Ok(0))
|| matches!(rejected_read, Err(ref error) if matches!(
error.kind(), io::ErrorKind::ConnectionReset | io::ErrorKind::ConnectionAborted
))
);
assert_eq!(budget.snapshot().in_flight, 1);
assert_eq!(budget.snapshot().rejected_total, 1);
drop((first, first_client, rejected_client));
let replacement_client = TcpStream::connect(second_listener.local_addr().unwrap())
.await
.unwrap();
let (replacement_io, _) = within(budget.accept(&second_listener)).await;
let replacement = budget.try_admit(replacement_io).unwrap();
assert_eq!(budget.snapshot().in_flight, 1);
assert_eq!(budget.snapshot().high_watermark, 1);
drop((replacement, replacement_client));
assert_eq!(budget.snapshot().in_flight, 0);
}
#[tokio::test]
async fn http_connection_budget_task_cancelled_before_first_poll_returns_permit() {
let (client, server) = tcp_pair().await;
let budget = Arc::new(HttpConnectionBudget::new(1));
let admitted = budget.try_admit(server).unwrap();
let polled = Arc::new(AtomicBool::new(false));
let polled_by_task = Arc::clone(&polled);
let task = tokio::spawn(async move {
let _io = admitted;
polled_by_task.store(true, Ordering::Relaxed);
std::future::pending::<()>().await;
});
task.abort();
assert!(within(task).await.unwrap_err().is_cancelled());
assert!(!polled.load(Ordering::Relaxed));
assert_eq!(budget.snapshot().in_flight, 0);
drop(client);
}
#[tokio::test]
async fn http_connection_budget_header_timeout_and_parse_failure_return_permit() {
for malformed in [false, true] {
let (mut client, server_io) = tcp_pair().await;
let budget = Arc::new(HttpConnectionBudget::new(1));
let admitted = budget.try_admit(server_io).unwrap();
let requests = Arc::new(AtomicUsize::new(0));
let requests_seen = Arc::clone(&requests);
let server = tokio::spawn(async move {
let service = service_fn(move |_: Request<Incoming>| {
requests_seen.fetch_add(1, Ordering::Relaxed);
async { Ok::<_, Infallible>(Response::new(Empty::<Bytes>::new())) }
});
hyper::server::conn::http1::Builder::new()
.timer(TokioTimer::new())
.header_read_timeout(Duration::from_millis(20))
.serve_connection(TokioIo::new(admitted), service)
.await
});
if malformed {
client.write_all(b"invalid request\r\n\r\n").await.unwrap();
}
let _connection_result = within(server).await.unwrap();
assert_eq!(requests.load(Ordering::Relaxed), 0);
assert_eq!(budget.snapshot().in_flight, 0);
drop(client);
}
}
#[tokio::test]
async fn http_connection_budget_h1_upgrade_keeps_permit_after_connection_future_finishes() {
let (mut client, server_io) = tcp_pair().await;
let budget = Arc::new(HttpConnectionBudget::new(1));
let admitted = budget.try_admit(server_io).unwrap();
let (upgrade_finished_tx, upgrade_finished_rx) = tokio::sync::oneshot::channel();
let upgrade_finished_tx = Arc::new(std::sync::Mutex::new(Some(upgrade_finished_tx)));
let server = tokio::spawn(async move {
let service = service_fn(move |mut request: Request<Incoming>| {
let on_upgrade = hyper::upgrade::on(&mut request);
let finished = upgrade_finished_tx.lock().unwrap().take().unwrap();
tokio::spawn(async move {
let mut upgraded = TokioIo::new(on_upgrade.await.unwrap());
let mut message = [0; 4];
upgraded.read_exact(&mut message).await.unwrap();
upgraded.write_all(&message).await.unwrap();
assert_eq!(upgraded.read(&mut [0; 1]).await.unwrap(), 0);
drop(upgraded);
let _ = finished.send(());
});
async {
Ok::<_, Infallible>(
Response::builder()
.status(StatusCode::SWITCHING_PROTOCOLS)
.header("connection", "upgrade")
.header("upgrade", "echo")
.body(Empty::<Bytes>::new())
.unwrap(),
)
}
});
hyper::server::conn::http1::Builder::new()
.serve_connection(TokioIo::new(admitted), service)
.with_upgrades()
.await
});
client
.write_all(
b"GET / HTTP/1.1\r\nHost: localhost\r\nConnection: upgrade\r\nUpgrade: echo\r\n\r\n",
)
.await
.unwrap();
let response = within(async {
let mut headers = Vec::new();
while !headers.ends_with(b"\r\n\r\n") {
assert!(headers.len() < 4096);
headers.push(client.read_u8().await.unwrap());
}
headers
})
.await;
assert!(response.starts_with(b"HTTP/1.1 101"));
within(server).await.unwrap().unwrap();
assert_eq!(budget.snapshot().in_flight, 1);
assert!(budget.try_admit(()).is_err());
client.write_all(b"ping").await.unwrap();
let mut echoed = [0; 4];
within(client.read_exact(&mut echoed)).await.unwrap();
assert_eq!(&echoed, b"ping");
drop(client);
within(upgrade_finished_rx).await.unwrap();
assert_eq!(budget.snapshot().in_flight, 0);
}
#[tokio::test]
async fn http_connection_budget_h2_parallel_streams_share_one_socket_permit() {
let (client_io, server_io) = tcp_pair().await;
let budget = Arc::new(HttpConnectionBudget::new(1));
let admitted = budget.try_admit(server_io).unwrap();
let requests = Arc::new(AtomicUsize::new(0));
let concurrent = Arc::new(tokio::sync::Barrier::new(3));
let server_requests = Arc::clone(&requests);
let server_concurrent = Arc::clone(&concurrent);
let server = tokio::spawn(async move {
let service = service_fn(move |_: Request<Incoming>| {
server_requests.fetch_add(1, Ordering::Relaxed);
let concurrent = Arc::clone(&server_concurrent);
async move {
concurrent.wait().await;
Ok::<_, Infallible>(Response::new(Empty::<Bytes>::new()))
}
});
hyper::server::conn::http2::Builder::new(TokioExecutor::new())
.max_concurrent_streams(2)
.serve_connection(TokioIo::new(admitted), service)
.await
});
let (sender, connection) = within(
hyper::client::conn::http2::Builder::new(TokioExecutor::new())
.handshake::<_, Empty<Bytes>>(TokioIo::new(client_io)),
)
.await
.unwrap();
let client_driver = tokio::spawn(connection);
let mut requests_in_flight = tokio::task::JoinSet::new();
for path in ["first", "second"] {
let mut sender = sender.clone();
requests_in_flight.spawn(async move {
let response = sender
.send_request(
Request::builder()
.uri(format!("http://localhost/{path}"))
.body(Empty::<Bytes>::new())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
response.into_body().collect().await.unwrap();
});
}
within(concurrent.wait()).await;
assert_eq!(requests.load(Ordering::Relaxed), 2);
assert_eq!(budget.snapshot().in_flight, 1);
assert_eq!(budget.snapshot().high_watermark, 1);
within(async {
while let Some(result) = requests_in_flight.join_next().await {
result.unwrap();
}
})
.await;
assert_eq!(budget.snapshot().in_flight, 1);
drop(sender);
client_driver.abort();
let _ = within(client_driver).await;
let _ = within(server).await.unwrap();
assert_eq!(budget.snapshot().in_flight, 0);
}
#[tokio::test(start_paused = true)]
async fn http_connection_accept_retries_peer_errors_and_backs_off_resource_errors() {
let budget = HttpConnectionBudget::new(1);
let mut attempts = VecDeque::from([
Err(io::Error::from(io::ErrorKind::ConnectionAborted)),
Err(io::Error::from(io::ErrorKind::ConnectionReset)),
Err(io::Error::other("injected file descriptor exhaustion")),
Err(io::Error::other("injected temporary accept failure")),
Ok(42),
]);
let started = tokio::time::Instant::now();
let accepted = budget
.accept_with(|| std::future::ready(attempts.pop_front().unwrap()))
.await;
assert_eq!(accepted, 42);
assert!(attempts.is_empty());
assert_eq!(started.elapsed(), Duration::from_secs(2));
assert_eq!(budget.snapshot().accept_errors_total, 4);
assert_eq!(budget.snapshot().in_flight, 0);
assert_eq!(budget.snapshot().rejected_total, 0);
}
+5 -1
View File
@@ -1,11 +1,15 @@
pub mod body;
mod connection;
pub mod middleware;
mod request_id;
pub use body::{
BodyBufferError, BodyBufferPolicy, BodyBufferReservation, BufferedBody,
BodyBufferBudget, BodyBufferError, BodyBufferPolicy, BodyBufferReservation, BufferedBody,
DEFAULT_BODY_BUFFER_PERMIT_BYTES,
};
pub use connection::{
http_connection_limit, AdmittedConnection, HttpConnectionBudget, HttpConnectionBudgetSnapshot,
};
pub use middleware::access_log::{
access_log_middleware, sanitize_access_log_path, should_downgrade_access_log,
+2 -1
View File
@@ -34,5 +34,6 @@ pub use queue::{
pub use redaction::{summarize_text_payload, TextPayloadSummary};
pub use shutdown::wait_for_shutdown_signal;
pub use tracing::{
init_reloadable_service_tracing, init_reloadable_tracing, LogFormat, LogReloader,
init_reloadable_service_tracing, init_reloadable_tracing, logging_metric_samples,
shutdown_logging, LogFormat, LogReloader, LogShutdownGuard,
};
+216 -60
View File
@@ -21,6 +21,11 @@ use crate::config::ServiceRuntimeConfig;
use crate::error::RuntimeBootstrapError;
use crate::observability::{FileLoggingConfig, LogDestination, LogRotation};
mod writer;
pub use writer::{logging_metric_samples, shutdown_logging, LogShutdownGuard};
use writer::{register_log_workers, LogWorker, NonBlockingLogWriter};
static TRACING_INIT: OnceLock<Result<(), String>> = OnceLock::new();
pub type LogReloader = Box<dyn Fn(&str) + Send + Sync>;
@@ -385,17 +390,12 @@ pub(crate) fn init_tracing(config: ServiceRuntimeConfig) -> Result<(), RuntimeBo
.unwrap_or_else(|_| config.default_log_filter.into());
let identity = RuntimeLogIdentity::from_config(&config);
let (file_writer, startup_cleanup_warning) =
if config.observability.log_destination.needs_file_sink() {
let Some(file_logging) = config.observability.file_logging.clone() else {
return Err("file logging requires a configured log directory".to_string());
};
let (writer, startup_cleanup_warning) =
RollingFileMakeWriter::new(config.service_name, file_logging)?;
(Some(writer), startup_cleanup_warning)
} else {
(None, None)
};
let RuntimeLogWriters {
stdout_writer,
file_writer,
workers,
startup_cleanup_warning,
} = RuntimeLogWriters::new(&config)?;
let init_result = match (
config.observability.log_format,
@@ -403,16 +403,26 @@ pub(crate) fn init_tracing(config: ServiceRuntimeConfig) -> Result<(), RuntimeBo
) {
(LogFormat::Pretty, LogDestination::Stdout) => tracing_subscriber::registry()
.with(filter)
.with(tracing_subscriber::fmt::layer().event_format(
PrettyRuntimeEventFormatter::new(identity.clone(), stdout_supports_ansi()),
))
.with(
tracing_subscriber::fmt::layer()
.event_format(PrettyRuntimeEventFormatter::new(
identity.clone(),
stdout_supports_ansi(),
))
.with_writer(
stdout_writer.clone().expect("stdout writer should exist"),
),
)
.try_init(),
(LogFormat::Json, LogDestination::Stdout) => tracing_subscriber::registry()
.with(filter)
.with(
tracing_subscriber::fmt::layer()
.json()
.event_format(JsonRuntimeEventFormatter::new(identity.clone())),
.event_format(JsonRuntimeEventFormatter::new(identity.clone()))
.with_writer(
stdout_writer.clone().expect("stdout writer should exist"),
),
)
.try_init(),
(LogFormat::Pretty, LogDestination::File) => tracing_subscriber::registry()
@@ -435,9 +445,16 @@ pub(crate) fn init_tracing(config: ServiceRuntimeConfig) -> Result<(), RuntimeBo
.try_init(),
(LogFormat::Pretty, LogDestination::Both) => tracing_subscriber::registry()
.with(filter)
.with(tracing_subscriber::fmt::layer().event_format(
PrettyRuntimeEventFormatter::new(identity.clone(), stdout_supports_ansi()),
))
.with(
tracing_subscriber::fmt::layer()
.event_format(PrettyRuntimeEventFormatter::new(
identity.clone(),
stdout_supports_ansi(),
))
.with_writer(
stdout_writer.clone().expect("stdout writer should exist"),
),
)
.with(
tracing_subscriber::fmt::layer()
.with_ansi(false)
@@ -450,7 +467,10 @@ pub(crate) fn init_tracing(config: ServiceRuntimeConfig) -> Result<(), RuntimeBo
.with(
tracing_subscriber::fmt::layer()
.json()
.event_format(JsonRuntimeEventFormatter::new(identity.clone())),
.event_format(JsonRuntimeEventFormatter::new(identity.clone()))
.with_writer(
stdout_writer.clone().expect("stdout writer should exist"),
),
)
.with(
tracing_subscriber::fmt::layer()
@@ -463,6 +483,7 @@ pub(crate) fn init_tracing(config: ServiceRuntimeConfig) -> Result<(), RuntimeBo
.map_err(|err| err.to_string());
if init_result.is_ok() {
register_log_workers(workers);
if let Some(warning) = startup_cleanup_warning.as_ref() {
emit_log_cleanup_warning("startup", warning.log_dir.as_path(), &warning.error);
}
@@ -497,39 +518,35 @@ pub fn init_reloadable_service_tracing(
let (filter_layer, reload_handle) = reload::Layer::new(filter);
let identity = RuntimeLogIdentity::from_config(&config);
let (file_writer, startup_cleanup_warning) =
if config.observability.log_destination.needs_file_sink() {
let Some(file_logging) = config.observability.file_logging.clone() else {
return Err(RuntimeBootstrapError::Tracing(
"file logging requires a configured log directory".to_string(),
));
};
let (writer, startup_cleanup_warning) =
RollingFileMakeWriter::new(config.service_name, file_logging)
.map_err(RuntimeBootstrapError::Tracing)?;
(Some(writer), startup_cleanup_warning)
} else {
(None, None)
};
let RuntimeLogWriters {
stdout_writer,
file_writer,
workers,
startup_cleanup_warning,
} = RuntimeLogWriters::new(&config).map_err(RuntimeBootstrapError::Tracing)?;
match (
config.observability.log_format,
config.observability.log_destination,
) {
(LogFormat::Pretty, LogDestination::Stdout) => {
tracing_subscriber::registry()
.with(filter_layer)
.with(tracing_subscriber::fmt::layer().event_format(
PrettyRuntimeEventFormatter::new(identity.clone(), stdout_supports_ansi()),
))
.try_init()
}
(LogFormat::Pretty, LogDestination::Stdout) => tracing_subscriber::registry()
.with(filter_layer)
.with(
tracing_subscriber::fmt::layer()
.event_format(PrettyRuntimeEventFormatter::new(
identity.clone(),
stdout_supports_ansi(),
))
.with_writer(stdout_writer.clone().expect("stdout writer should exist")),
)
.try_init(),
(LogFormat::Json, LogDestination::Stdout) => tracing_subscriber::registry()
.with(filter_layer)
.with(
tracing_subscriber::fmt::layer()
.json()
.event_format(JsonRuntimeEventFormatter::new(identity.clone())),
.event_format(JsonRuntimeEventFormatter::new(identity.clone()))
.with_writer(stdout_writer.clone().expect("stdout writer should exist")),
)
.try_init(),
(LogFormat::Pretty, LogDestination::File) => tracing_subscriber::registry()
@@ -550,26 +567,30 @@ pub fn init_reloadable_service_tracing(
.with_writer(file_writer.clone().expect("file writer should exist")),
)
.try_init(),
(LogFormat::Pretty, LogDestination::Both) => {
tracing_subscriber::registry()
.with(filter_layer)
.with(tracing_subscriber::fmt::layer().event_format(
PrettyRuntimeEventFormatter::new(identity.clone(), stdout_supports_ansi()),
))
.with(
tracing_subscriber::fmt::layer()
.with_ansi(false)
.event_format(PrettyRuntimeEventFormatter::new(identity.clone(), false))
.with_writer(file_writer.clone().expect("file writer should exist")),
)
.try_init()
}
(LogFormat::Pretty, LogDestination::Both) => tracing_subscriber::registry()
.with(filter_layer)
.with(
tracing_subscriber::fmt::layer()
.event_format(PrettyRuntimeEventFormatter::new(
identity.clone(),
stdout_supports_ansi(),
))
.with_writer(stdout_writer.clone().expect("stdout writer should exist")),
)
.with(
tracing_subscriber::fmt::layer()
.with_ansi(false)
.event_format(PrettyRuntimeEventFormatter::new(identity.clone(), false))
.with_writer(file_writer.clone().expect("file writer should exist")),
)
.try_init(),
(LogFormat::Json, LogDestination::Both) => tracing_subscriber::registry()
.with(filter_layer)
.with(
tracing_subscriber::fmt::layer()
.json()
.event_format(JsonRuntimeEventFormatter::new(identity.clone())),
.event_format(JsonRuntimeEventFormatter::new(identity.clone()))
.with_writer(stdout_writer.clone().expect("stdout writer should exist")),
)
.with(
tracing_subscriber::fmt::layer()
@@ -581,6 +602,7 @@ pub fn init_reloadable_service_tracing(
}
.map_err(|err| RuntimeBootstrapError::Tracing(err.to_string()))?;
register_log_workers(workers);
if let Some(warning) = startup_cleanup_warning.as_ref() {
emit_log_cleanup_warning("startup", warning.log_dir.as_path(), &warning.error);
}
@@ -595,6 +617,54 @@ pub fn init_reloadable_service_tracing(
}))
}
struct RuntimeLogWriters {
stdout_writer: Option<NonBlockingLogWriter>,
file_writer: Option<NonBlockingLogWriter>,
workers: Vec<LogWorker>,
startup_cleanup_warning: Option<StartupCleanupWarning>,
}
impl RuntimeLogWriters {
fn new(config: &ServiceRuntimeConfig) -> Result<Self, String> {
let (file_sink, startup_cleanup_warning) =
if config.observability.log_destination.needs_file_sink() {
let file_logging = config
.observability
.file_logging
.clone()
.ok_or("file logging requires a configured log directory")?;
let (sink, warning) =
RollingFileMakeWriter::new(config.service_name, file_logging)?;
(Some(sink), warning)
} else {
(None, None)
};
let mut workers = Vec::with_capacity(2);
let stdout_writer = if config.observability.log_destination != LogDestination::File {
let (writer, worker) = NonBlockingLogWriter::new("stdout", io::stdout())
.map_err(|err| format!("failed to start stdout log writer: {err}"))?;
workers.push(worker);
Some(writer)
} else {
None
};
let file_writer = if let Some(sink) = file_sink {
let (writer, worker) = NonBlockingLogWriter::new("file", sink.make_writer())
.map_err(|err| format!("failed to start file log writer: {err}"))?;
workers.push(worker);
Some(writer)
} else {
None
};
Ok(Self {
stdout_writer,
file_writer,
workers,
startup_cleanup_warning,
})
}
}
#[derive(Debug, Clone)]
struct RollingFileMakeWriter {
sink: Arc<RollingFileSink>,
@@ -692,7 +762,10 @@ impl RollingFileSink {
}
fn write(&self, buf: &[u8]) -> io::Result<usize> {
let now = Local::now();
self.write_at(buf, Local::now())
}
fn write_at(&self, buf: &[u8], now: DateTime<Local>) -> io::Result<usize> {
let mut state = self
.state
.lock()
@@ -833,8 +906,19 @@ fn spawn_log_cleanup_task(service_name: &'static str, config: FileLoggingConfig)
let interval = Duration::from_secs(6 * 60 * 60);
loop {
tokio::time::sleep(interval).await;
if let Err(err) = cleanup_log_files(service_name, &config) {
emit_log_cleanup_warning("background", config.dir.as_path(), &err);
let cleanup_config = config.clone();
match tokio::task::spawn_blocking(move || {
cleanup_log_files(service_name, &cleanup_config)
})
.await
{
Ok(Ok(_)) => {}
Ok(Err(err)) => {
emit_log_cleanup_warning("background", config.dir.as_path(), &err);
}
Err(err) => {
emit_log_cleanup_warning("background", config.dir.as_path(), &err);
}
}
}
});
@@ -1028,6 +1112,78 @@ mod tests {
);
}
#[test]
fn rolling_file_sink_rotates_without_mixing_bucket_contents() {
for rotation in [LogRotation::Hourly, LogRotation::Daily] {
let dir = std::env::temp_dir().join(format!("aether-runtime-logs-{}", Uuid::new_v4()));
let config = FileLoggingConfig::new(&dir, rotation, 7, 30);
let (sink, _) = RollingFileSink::new("runtime-test", config).expect("sink should open");
let before = Local
.with_ymd_and_hms(2026, 4, 4, 23, 59, 59)
.single()
.expect("timestamp should build");
let after = before + chrono::Duration::seconds(2);
assert_eq!(sink.write_at(b"before\n", before).unwrap(), 7);
assert_eq!(sink.write_at(b"after\n", after).unwrap(), 6);
sink.flush().unwrap();
for (instant, expected) in [(before, "before\n"), (after, "after\n")] {
let path =
bucketed_log_path(&dir, "runtime-test", &log_bucket_key(rotation, instant));
assert_eq!(fs::read_to_string(&path).unwrap(), expected);
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt as _;
assert_eq!(
fs::metadata(path).unwrap().permissions().mode() & 0o777,
0o600
);
}
}
drop(sink);
fs::remove_dir_all(&dir).unwrap();
}
}
#[cfg(unix)]
#[test]
fn failed_rotation_preserves_old_file_and_can_retry_safely() {
use std::os::unix::fs::symlink;
let dir = std::env::temp_dir().join(format!("aether-runtime-logs-{}", Uuid::new_v4()));
let config = FileLoggingConfig::new(&dir, LogRotation::Daily, 7, 30);
let (sink, _) = RollingFileSink::new("runtime-test", config).unwrap();
let before = Local
.with_ymd_and_hms(2026, 4, 4, 12, 0, 0)
.single()
.unwrap();
let after = before + chrono::Duration::days(1);
let old_bucket = log_bucket_key(LogRotation::Daily, before);
let old_path = bucketed_log_path(&dir, "runtime-test", &old_bucket);
let new_path = bucketed_log_path(
&dir,
"runtime-test",
&log_bucket_key(LogRotation::Daily, after),
);
let victim = dir.join("victim.txt");
fs::write(&victim, b"unchanged").unwrap();
sink.write_at(b"before\n", before).unwrap();
symlink(&victim, &new_path).unwrap();
assert!(sink.write_at(b"rejected\n", after).is_err());
assert_eq!(sink.state.lock().unwrap().current_bucket, old_bucket);
assert_eq!(fs::read(&victim).unwrap(), b"unchanged");
assert_eq!(fs::read(&old_path).unwrap(), b"before\n");
fs::remove_file(&new_path).unwrap();
sink.write_at(b"after\n", after).unwrap();
sink.flush().unwrap();
assert_eq!(fs::read(&new_path).unwrap(), b"after\n");
assert_eq!(fs::read(&old_path).unwrap(), b"before\n");
drop(sink);
fs::remove_dir_all(&dir).unwrap();
}
#[test]
fn cleanup_log_files_removes_matching_files_on_disk() {
let dir = std::env::temp_dir().join(format!("aether-runtime-logs-{}", Uuid::new_v4()));
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,253 @@
use std::fs;
use std::io::{self, Read};
use std::path::{Path, PathBuf};
use std::process::{Child, Command, ExitStatus, Stdio};
use std::thread::JoinHandle;
use std::time::{Duration, Instant};
use aether_runtime::{
init_service_runtime, logging_metric_samples, FileLoggingConfig, LogDestination, LogFormat,
LogRotation, LogShutdownGuard, ServiceRuntimeConfig,
};
const CHILD_ENV: &str = "AETHER_TEST_BLOCKED_STDOUT_CHILD";
const DIRECTORY_ENV: &str = "AETHER_TEST_BLOCKED_STDOUT_DIR";
const SERVICE_NAME: &str = "blocked-stdout-test";
const EVENT_NAME: &str = "blocked_stdout_probe";
const TEST_NAME: &str = "blocked_stdout_does_not_block_file_logs_or_process_exit";
#[test]
fn blocked_stdout_does_not_block_file_logs_or_process_exit() {
if std::env::var_os(CHILD_ENV).is_some() {
run_child_scenario();
eprintln!("blocked stdout guard returned");
// The scenario returns normally and drops its guard. Skip libtest's own
// stdout report, while still exercising Rust's standard exit cleanup.
std::process::exit(0);
}
let directory = TestDirectory::new();
let mut command = Command::new(std::env::current_exe().expect("test executable"));
command
.args(["--exact", TEST_NAME, "--nocapture", "--quiet"])
.env(CHILD_ENV, "1")
.env(DIRECTORY_ENV, &directory.0)
.env_remove("RUST_LOG")
.env_remove("NO_COLOR")
.env_remove("FORCE_COLOR");
let mut child = BlockedStdoutChild::spawn(&mut command).expect("logging child should start");
let (status, timed_out, stderr) = child
.wait(Duration::from_secs(8))
.expect("logging child should be reaped");
let stderr = String::from_utf8_lossy(&stderr);
assert!(
!timed_out,
"blocked stdout prevented process exit within 8 seconds: {stderr}"
);
assert!(status.success(), "logging child failed: {status}: {stderr}");
assert!(
stderr.contains("blocked stdout saturated")
&& stderr.contains("blocked stdout guard returned"),
"child did not reach saturation and return from its guard: {stderr}"
);
let records = read_file_records(&directory.0);
assert!(
!records.is_empty(),
"healthy file destination received no logs"
);
let final_markers: Vec<_> = records
.iter()
.filter(|record| record["fields"]["phase"] == "final")
.collect();
assert_eq!(
final_markers.len(),
1,
"file lost or duplicated final marker"
);
let final_marker = final_markers[0];
assert_eq!(final_marker["fields"]["event_name"], EVENT_NAME);
assert!(
final_marker["fields"]["stdout_dropped_full"]
.as_u64()
.expect("stdout queue drop counter")
+ final_marker["fields"]["stdout_dropped_bytes"]
.as_u64()
.expect("stdout byte drop counter")
> 0,
"file marker must prove stdout saturation"
);
assert_eq!(records.last().unwrap()["fields"]["phase"], "final");
}
fn run_child_scenario() {
let _shutdown = LogShutdownGuard::new();
let directory = PathBuf::from(std::env::var_os(DIRECTORY_ENV).expect("child log directory"));
init_service_runtime(
ServiceRuntimeConfig::new(SERVICE_NAME, "info")
.with_log_destination(LogDestination::Both)
.with_log_format(LogFormat::Json)
.with_file_logging(FileLoggingConfig::new(directory, LogRotation::Daily, 7, 30)),
)
.expect("both log destinations should initialize");
let payload = "x".repeat(384);
let flood_deadline = Instant::now() + Duration::from_secs(3);
let mut emitted = 0;
while emitted < 20_000 && Instant::now() < flood_deadline {
tracing::info!(
event_name = EVENT_NAME,
phase = "flood",
sequence = emitted,
payload = %payload,
"fill unread stdout"
);
emitted += 1;
if emitted % 64 == 0 && stdout_dropped_events() > 0 {
break;
}
}
assert!(stdout_dropped_events() > 0, "stdout queue did not saturate");
let file_deadline = Instant::now() + Duration::from_secs(1);
while metric("logging_file_retained_bytes") > 0 && Instant::now() < file_deadline {
std::thread::sleep(Duration::from_millis(5));
}
assert_eq!(
metric("logging_file_retained_bytes"),
0,
"file writer stalled"
);
assert_eq!(metric("logging_file_write_errors_total"), 0);
let dropped_full = metric("logging_stdout_dropped_full_total");
let dropped_bytes = metric("logging_stdout_dropped_bytes_total");
eprintln!(
"blocked stdout saturated: full={dropped_full} bytes={dropped_bytes} emitted={emitted}"
);
tracing::info!(
event_name = EVENT_NAME,
phase = "final",
stdout_dropped_full = dropped_full,
stdout_dropped_bytes = dropped_bytes,
"healthy file final marker"
);
}
fn metric(name: &str) -> u64 {
logging_metric_samples()
.into_iter()
.find(|sample| sample.name == name)
.unwrap_or_else(|| panic!("missing logging metric: {name}"))
.value
}
fn stdout_dropped_events() -> u64 {
metric("logging_stdout_dropped_full_total") + metric("logging_stdout_dropped_bytes_total")
}
fn read_file_records(directory: &Path) -> Vec<serde_json::Value> {
let mut paths: Vec<_> = fs::read_dir(directory)
.expect("log directory")
.map(|entry| entry.expect("log entry").path())
.filter(|path| {
path.file_name()
.and_then(|name| name.to_str())
.is_some_and(|name| name.starts_with(SERVICE_NAME) && name.ends_with(".log"))
})
.collect();
paths.sort();
let mut records = Vec::new();
for path in paths {
let contents = fs::read_to_string(path).expect("UTF-8 file logs");
assert!(contents.ends_with('\n'), "partial final file record");
for line in contents.lines() {
assert!(line.len() < 1024, "test event exceeded 1 KiB");
records.push(serde_json::from_str(line).expect("complete JSON file record"));
}
}
records
}
struct BlockedStdoutChild {
child: Child,
stderr_reader: Option<JoinHandle<io::Result<Vec<u8>>>>,
}
impl BlockedStdoutChild {
fn spawn(command: &mut Command) -> io::Result<Self> {
let child = command
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn()?;
let mut guarded = Self {
child,
stderr_reader: None,
};
let mut stderr = guarded.child.stderr.take().expect("child stderr pipe");
guarded.stderr_reader = Some(std::thread::Builder::new().spawn(move || {
let mut captured = Vec::new();
let mut chunk = [0u8; 1024];
loop {
let count = stderr.read(&mut chunk)?;
if count == 0 {
return Ok(captured);
}
let retained = count.min((16 * 1024usize).saturating_sub(captured.len()));
captured.extend_from_slice(&chunk[..retained]);
}
})?);
Ok(guarded)
}
fn wait(&mut self, timeout: Duration) -> io::Result<(ExitStatus, bool, Vec<u8>)> {
let deadline = Instant::now() + timeout;
let (status, timed_out) = loop {
if let Some(status) = self.child.try_wait()? {
break (status, false);
}
if Instant::now() >= deadline {
let _ = self.child.kill();
break (self.child.wait()?, true);
}
std::thread::sleep(Duration::from_millis(10));
};
// Keep the stdout read end open and completely unread until the child
// has exited or been killed. Closing it earlier would unblock writes.
drop(self.child.stdout.take());
let stderr = self
.stderr_reader
.take()
.expect("stderr reader")
.join()
.map_err(|_| io::Error::other("stderr reader panicked"))??;
Ok((status, timed_out, stderr))
}
}
impl Drop for BlockedStdoutChild {
fn drop(&mut self) {
let _ = self.child.kill();
let _ = self.child.wait();
drop(self.child.stdout.take());
if let Some(reader) = self.stderr_reader.take() {
let _ = reader.join();
}
}
}
struct TestDirectory(PathBuf);
impl TestDirectory {
fn new() -> Self {
let path =
std::env::temp_dir().join(format!("aether-blocked-stdout-{}", uuid::Uuid::new_v4()));
fs::create_dir(&path).expect("test directory");
Self(path)
}
}
impl Drop for TestDirectory {
fn drop(&mut self) {
let _ = fs::remove_dir_all(&self.0);
}
}
@@ -0,0 +1,354 @@
use std::fs;
use std::io::Read;
use std::path::{Path, PathBuf};
use std::process::{Command, Output, Stdio};
use std::time::{Duration, Instant};
use aether_runtime::{
init_reloadable_service_tracing, init_service_runtime, FileLoggingConfig, LogDestination,
LogFormat, LogRotation, LogShutdownGuard, ServiceRuntimeConfig,
};
const CASE_ENV: &str = "AETHER_TEST_NONBLOCKING_LOGGING_CASE";
const DIRECTORY_ENV: &str = "AETHER_TEST_NONBLOCKING_LOGGING_DIR";
const EVENT_NAME: &str = "nonblocking_logging_probe";
const SERVICE_NAME: &str = "nonblocking-logging-test";
const RECORD_COUNT: usize = 32;
const FIELD_VALUE: &str = "quote\" newline\n backslash\\ \u{4e2d}\u{6587}";
#[test]
fn nonblocking_logging_entrypoints_reload_and_guard_drain() {
if let Ok(scenario) = std::env::var(CASE_ENV) {
run_scenario(&scenario);
return;
}
let root = TestDirectory::new();
for entrypoint in ["standard", "reloadable"] {
for destination in ["stdout", "file", "both"] {
for format in ["pretty", "json"] {
let scenario = format!("{entrypoint}-{destination}-{format}");
let directory = root.0.join(&scenario);
fs::create_dir(&directory).expect("scenario directory");
let mut command = Command::new(std::env::current_exe().expect("test executable"));
command
.args([
"--exact",
"nonblocking_logging_entrypoints_reload_and_guard_drain",
"--nocapture",
"--quiet",
])
.env(CASE_ENV, &scenario)
.env(DIRECTORY_ENV, &directory)
.env_remove("RUST_LOG")
.env_remove("NO_COLOR")
.env_remove("FORCE_COLOR");
let output = run_subprocess(&mut command);
let stdout = String::from_utf8(output.stdout).expect("UTF-8 stdout");
let stderr = String::from_utf8(output.stderr).expect("UTF-8 stderr");
assert!(output.status.success(), "{scenario}: {stdout}\n{stderr}");
let file_output = read_log_files(&directory);
verify_output(
&stdout,
entrypoint,
format,
destination != "file",
&scenario,
);
verify_output(
&file_output,
entrypoint,
format,
destination != "stdout",
&scenario,
);
}
}
}
}
fn run_scenario(scenario: &str) {
let parts: Vec<_> = scenario.split('-').collect();
let [entrypoint, destination, format] = parts.as_slice() else {
panic!("invalid logging scenario: {scenario}");
};
let _shutdown = LogShutdownGuard::new();
let destination = match *destination {
"stdout" => LogDestination::Stdout,
"file" => LogDestination::File,
"both" => LogDestination::Both,
other => panic!("unknown destination: {other}"),
};
let mut config = ServiceRuntimeConfig::new(SERVICE_NAME, "info")
.with_node_role("integration")
.with_instance_id("logging-child")
.with_log_destination(destination)
.with_log_format(match *format {
"pretty" => LogFormat::Pretty,
"json" => LogFormat::Json,
other => panic!("unknown format: {other}"),
});
if matches!(destination, LogDestination::File | LogDestination::Both) {
config = config.with_file_logging(FileLoggingConfig::new(
PathBuf::from(std::env::var_os(DIRECTORY_ENV).expect("scenario log directory")),
LogRotation::Daily,
7,
30,
));
}
let reload = match *entrypoint {
"standard" => {
init_service_runtime(config).expect("standard logging initializes");
None
}
"reloadable" => Some(
init_reloadable_service_tracing("info", config)
.expect("reloadable logging initializes"),
),
other => panic!("unknown entrypoint: {other}"),
};
tracing::debug!(
event_name = EVENT_NAME,
phase = "initial_hidden",
"filtered debug"
);
tracing::info!(event_name = EVENT_NAME, phase = "initial", "initial event");
if let Some(reload) = reload {
reload("debug");
tracing::debug!(
event_name = EVENT_NAME,
phase = "reloaded_debug",
"visible debug"
);
let invalid_filter = "nonblocking_logging=not-a-level";
assert!(tracing_subscriber::EnvFilter::try_new(invalid_filter).is_err());
reload(invalid_filter);
tracing::debug!(
event_name = EVENT_NAME,
phase = "invalid_reload_unchanged",
"still debug"
);
reload("error");
tracing::info!(
event_name = EVENT_NAME,
phase = "error_filter_hidden",
"filtered info"
);
tracing::error!(
event_name = EVENT_NAME,
phase = "reloaded_error",
"visible error"
);
reload("info");
}
for sequence in 0..RECORD_COUNT {
tracing::info!(
event_name = EVENT_NAME,
phase = "record",
sequence = sequence as u64,
value = FIELD_VALUE,
"complete record"
);
}
tracing::info!(
event_name = EVENT_NAME,
phase = "tail",
"final event before guard drop"
);
// Returning drops the guard. The parent verifies the tail after process exit.
}
fn verify_output(output: &str, entrypoint: &str, format: &str, enabled: bool, scenario: &str) {
let lines: Vec<_> = output
.lines()
.filter(|line| line.contains(EVENT_NAME))
.collect();
if !enabled {
assert!(
lines.is_empty(),
"unexpected destination output in {scenario}: {output}"
);
return;
}
let expected_count = RECORD_COUNT + 2 + usize::from(entrypoint == "reloadable") * 3;
assert_eq!(
lines.len(),
expected_count,
"missing or duplicate records in {scenario}: {output}"
);
assert!(
!output.contains('\u{1b}'),
"redirected/file output must not contain ANSI: {scenario}"
);
assert!(
!output.contains("initial_hidden"),
"initial filter failed: {scenario}"
);
assert!(
!output.contains("error_filter_hidden"),
"reloaded filter failed: {scenario}"
);
let mut phases = Vec::new();
let mut sequences = Vec::new();
for line in lines {
if format == "json" {
let record: serde_json::Value = serde_json::from_str(line)
.unwrap_or_else(|error| panic!("incomplete JSON in {scenario}: {error}: {line}"));
assert_eq!(record["service"], SERVICE_NAME);
assert_eq!(record["node_role"], "integration");
assert_eq!(record["instance_id"], "logging-child");
let phase = record["fields"]["phase"].as_str().expect("event phase");
phases.push(phase.to_string());
if phase == "record" {
assert_eq!(record["fields"]["value"], FIELD_VALUE);
sequences.push(
record["fields"]["sequence"]
.as_u64()
.expect("record sequence"),
);
}
} else {
assert!(
line.contains(" | INFO") || line.contains(" | DEBUG") || line.contains(" | ERROR"),
"incomplete Pretty record in {scenario}: {line}"
);
let phase = [
"initial",
"reloaded_debug",
"invalid_reload_unchanged",
"reloaded_error",
"record",
"tail",
]
.into_iter()
.find(|phase| line.contains(&format!("phase=\"{phase}\"")))
.expect("complete Pretty phase field");
phases.push(phase.to_string());
if phase == "record" {
let expected_value = format!("value={FIELD_VALUE:?}");
assert!(
line.contains(&expected_value),
"incomplete Pretty value in {scenario}: {line}"
);
let sequence = line
.split_whitespace()
.find_map(|field| field.strip_prefix("sequence="))
.expect("complete Pretty sequence field")
.parse::<u64>()
.expect("sequence number");
sequences.push(sequence);
}
}
}
assert_eq!(phases.first().map(String::as_str), Some("initial"));
assert_eq!(
phases.last().map(String::as_str),
Some("tail"),
"guard lost tail event: {scenario}"
);
for phase in [
"reloaded_debug",
"invalid_reload_unchanged",
"reloaded_error",
] {
assert_eq!(
phases
.iter()
.filter(|value| value.as_str() == phase)
.count(),
usize::from(entrypoint == "reloadable"),
"reload phase {phase} in {scenario}"
);
}
assert_eq!(
sequences,
(0..RECORD_COUNT as u64).collect::<Vec<_>>(),
"records must remain complete and ordered in {scenario}"
);
}
fn read_log_files(directory: &Path) -> String {
let mut paths: Vec<_> = fs::read_dir(directory)
.expect("log directory")
.map(|entry| entry.expect("log directory entry").path())
.filter(|path| {
path.file_name()
.and_then(|name| name.to_str())
.is_some_and(|name| name.starts_with(SERVICE_NAME) && name.ends_with(".log"))
})
.collect();
paths.sort();
paths
.into_iter()
.map(|path| fs::read_to_string(path).expect("UTF-8 file log"))
.collect()
}
fn run_subprocess(command: &mut Command) -> Output {
let mut child = command
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn()
.expect("logging subprocess");
let stdout = child.stdout.take().expect("stdout pipe");
let stderr = child.stderr.take().expect("stderr pipe");
let stdout_reader = std::thread::spawn(move || read_pipe(stdout));
let stderr_reader = std::thread::spawn(move || read_pipe(stderr));
let deadline = Instant::now() + Duration::from_secs(10);
let mut timed_out = false;
let status = loop {
if let Some(status) = child.try_wait().expect("child status") {
break status;
}
if Instant::now() >= deadline {
timed_out = true;
let _ = child.kill();
break child.wait().expect("reap timed out child");
}
std::thread::sleep(Duration::from_millis(10));
};
let output = Output {
status,
stdout: stdout_reader.join().expect("stdout reader"),
stderr: stderr_reader.join().expect("stderr reader"),
};
assert!(
!timed_out,
"logging subprocess timed out: {}\n{}",
String::from_utf8_lossy(&output.stdout),
String::from_utf8_lossy(&output.stderr)
);
output
}
fn read_pipe(mut pipe: impl Read) -> Vec<u8> {
const MAX_CAPTURE_BYTES: usize = 4 * 1024 * 1024;
let mut captured = Vec::new();
let mut chunk = [0u8; 8192];
loop {
let count = pipe.read(&mut chunk).expect("drain child pipe");
if count == 0 {
return captured;
}
let retained = count.min(MAX_CAPTURE_BYTES.saturating_sub(captured.len()));
captured.extend_from_slice(&chunk[..retained]);
}
}
struct TestDirectory(PathBuf);
impl TestDirectory {
fn new() -> Self {
let path =
std::env::temp_dir().join(format!("aether-nonblocking-logs-{}", uuid::Uuid::new_v4()));
fs::create_dir(&path).expect("test directory");
Self(path)
}
}
impl Drop for TestDirectory {
fn drop(&mut self) {
let _ = fs::remove_dir_all(&self.0);
}
}
@@ -1,13 +1,15 @@
#![cfg(target_os = "linux")]
use std::fs;
use std::io::Read;
use std::os::unix::fs::{MetadataExt as _, PermissionsExt as _};
use std::path::PathBuf;
use std::process::Command;
use std::process::{Command, Output, Stdio};
use std::time::{Duration, Instant};
use aether_runtime::{
init_reloadable_service_tracing, init_service_runtime, FileLoggingConfig, LogDestination,
LogFormat, LogRotation, ServiceRuntimeConfig,
init_reloadable_service_tracing, init_service_runtime, shutdown_logging, FileLoggingConfig,
LogDestination, LogFormat, LogRotation, ServiceRuntimeConfig,
};
#[test]
@@ -23,7 +25,9 @@ fn root_appends_to_existing_logs_without_changing_ownership() {
for format in ["pretty", "json"] {
for owner in ["0", "1000", "65532", "new"] {
let scenario = format!("{entrypoint}-{destination}-{format}-{owner}");
let output = Command::new(std::env::current_exe().expect("test executable"))
let mut command =
Command::new(std::env::current_exe().expect("test executable"));
command
.args([
"--ignored",
"--exact",
@@ -31,9 +35,8 @@ fn root_appends_to_existing_logs_without_changing_ownership() {
"--nocapture",
])
.env("AETHER_TEST_ROOT_LOGGING_CASE", &scenario)
.env_remove("RUST_LOG")
.output()
.expect("root logging subprocess");
.env_remove("RUST_LOG");
let output = run_subprocess(&mut command);
let stdout = String::from_utf8_lossy(&output.stdout);
let stderr = String::from_utf8_lossy(&output.stderr);
assert!(output.status.success(), "{scenario}: {stdout}\n{stderr}");
@@ -49,6 +52,51 @@ fn root_appends_to_existing_logs_without_changing_ownership() {
}
}
fn run_subprocess(command: &mut Command) -> Output {
let mut child = command
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn()
.expect("root logging subprocess");
let mut stdout = child.stdout.take().expect("stdout pipe");
let mut stderr = child.stderr.take().expect("stderr pipe");
let stdout_reader = std::thread::spawn(move || {
let mut bytes = Vec::new();
stdout.read_to_end(&mut bytes).expect("read child stdout");
bytes
});
let stderr_reader = std::thread::spawn(move || {
let mut bytes = Vec::new();
stderr.read_to_end(&mut bytes).expect("read child stderr");
bytes
});
let deadline = Instant::now() + Duration::from_secs(10);
let mut timed_out = false;
let status = loop {
if let Some(status) = child.try_wait().expect("child status") {
break status;
}
if Instant::now() >= deadline {
timed_out = true;
let _ = child.kill();
break child.wait().expect("reap timed out child");
}
std::thread::sleep(Duration::from_millis(10));
};
let output = Output {
status,
stdout: stdout_reader.join().expect("stdout reader"),
stderr: stderr_reader.join().expect("stderr reader"),
};
assert!(
!timed_out,
"root logging subprocess timed out: {}\n{}",
String::from_utf8_lossy(&output.stdout),
String::from_utf8_lossy(&output.stderr)
);
output
}
fn run_scenario(scenario: &str) {
assert_eq!(unsafe { libc::geteuid() }, 0);
assert_eq!(unsafe { libc::getegid() }, 0);
@@ -117,6 +165,10 @@ fn run_scenario(scenario: &str) {
other => panic!("unknown entrypoint: {other}"),
};
tracing::info!("root logging ready");
assert!(
shutdown_logging(Duration::from_secs(2)),
"root file logging should drain"
);
let metadata = fs::metadata(&log_file).expect("written log file");
assert_eq!(metadata.uid(), expected_owner);
+598 -20
View File
@@ -1,6 +1,7 @@
mod error;
mod memory;
pub mod redis;
mod score_window;
use std::collections::BTreeMap;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
@@ -17,6 +18,7 @@ use async_trait::async_trait;
pub use error::DataLayerError;
use memory::MemoryRuntimeBackend;
pub use memory::MemoryRuntimeStateConfig;
pub use score_window::{ScoreWindowU64Stats, SCORE_WINDOW_AGGREGATION_MEMBER_LIMIT};
use tokio::task::JoinHandle;
use tracing::warn;
use uuid::Uuid;
@@ -621,6 +623,32 @@ impl RuntimeState {
}
}
/// Aggregate at most 512 timestamped `prefix:u64` members per key without
/// transferring their history. `None` requires an exact full-range fallback;
/// it never represents an empty or cached window.
pub async fn score_window_u64_stats_by_min(
&self,
keys: &[String],
min_score: f64,
) -> Result<Vec<Option<ScoreWindowU64Stats>>, DataLayerError> {
if !min_score.is_finite() {
return Err(DataLayerError::InvalidInput(
"runtime window minimum score must be finite".to_string(),
));
}
match self.backend.as_ref() {
RuntimeStateBackend::Memory(memory) => {
Ok(memory.score_window_u64_stats_by_min(keys, min_score).await)
}
RuntimeStateBackend::Redis(redis) => {
redis
.runtime
.score_window_u64_stats_by_min(keys, min_score)
.await
}
}
}
pub async fn score_remove_by_score(
&self,
key: &str,
@@ -1046,6 +1074,26 @@ pub struct RuntimeQueueEntry {
pub fields: BTreeMap<String, String>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RuntimeQueueReclaimPage {
/// Resume the next reclaim scan here; `0-0` marks the end of the current scan.
pub next_start_id: String,
pub entries: Vec<RuntimeQueueEntry>,
pub deleted_ids: Vec<String>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RuntimeQueueTransferOutcome {
Transferred {
destination_id: String,
acked: usize,
deleted: usize,
},
/// No pending entry was present. This does not assert that it was archived:
/// another consumer, deletion, or retention policy may have removed it.
NotPending,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct RuntimeQueueStats {
pub stream_length: u64,
@@ -1085,6 +1133,45 @@ fn validate_runtime_queue_reclaim_config(
Ok(())
}
pub(crate) fn validate_runtime_queue_transfer(
source: &str,
group: &str,
entry_id: &str,
destination: &str,
destination_fields: &BTreeMap<String, String>,
) -> Result<(), DataLayerError> {
validate_runtime_queue_name(source, "runtime queue source stream")?;
validate_runtime_queue_name(group, "runtime queue group")?;
validate_runtime_queue_name(destination, "runtime queue destination stream")?;
if source == destination {
return Err(DataLayerError::InvalidInput(
"runtime queue transfer source and destination must differ".to_string(),
));
}
if destination_fields.is_empty() {
return Err(DataLayerError::InvalidInput(
"runtime queue transfer destination fields cannot be empty".to_string(),
));
}
let canonical_u64 = |value: &str| {
!value.is_empty()
&& value.bytes().all(|byte| byte.is_ascii_digit())
&& (value.len() == 1 || !value.starts_with('0'))
&& value.parse::<u64>().is_ok()
};
if !entry_id
.split_once('-')
.is_some_and(|(milliseconds, sequence)| {
canonical_u64(milliseconds) && canonical_u64(sequence)
})
{
return Err(DataLayerError::InvalidInput(
"runtime queue transfer entry id must be a canonical u64-u64 stream id".to_string(),
));
}
Ok(())
}
#[async_trait]
pub trait RuntimeQueueStore: Send + Sync {
async fn ensure_consumer_group(
@@ -1119,6 +1206,39 @@ pub trait RuntimeQueueStore: Send + Sync {
config: RuntimeQueueReclaimConfig,
) -> Result<Vec<RuntimeQueueEntry>, DataLayerError>;
/// Existing queue backends can retain their complete-scan behavior without implementing paging.
async fn claim_stale_page(
&self,
stream: &str,
group: &str,
consumer: &str,
start_id: &str,
config: RuntimeQueueReclaimConfig,
) -> Result<RuntimeQueueReclaimPage, DataLayerError> {
Ok(RuntimeQueueReclaimPage {
next_start_id: "0-0".to_string(),
entries: self
.claim_stale(stream, group, consumer, start_id, config)
.await?,
deleted_ids: Vec::new(),
})
}
/// Atomically append caller-supplied fields, acknowledge the pending source entry, and
/// delete that source ID. Repeated calls must not append when the entry is no longer pending.
/// `None` means unsupported and has no side effects; callers may explicitly retain their
/// existing non-atomic fallback for third-party queue implementations.
async fn try_transfer_pending_to_stream(
&self,
_source: &str,
_group: &str,
_entry_id: &str,
_destination: &str,
_destination_fields: &BTreeMap<String, String>,
) -> Result<Option<RuntimeQueueTransferOutcome>, DataLayerError> {
Ok(None)
}
async fn ack(&self, stream: &str, group: &str, ids: &[String])
-> Result<usize, DataLayerError>;
@@ -1242,6 +1362,20 @@ impl RuntimeQueueStore for RuntimeState {
start_id: &str,
config: RuntimeQueueReclaimConfig,
) -> Result<Vec<RuntimeQueueEntry>, DataLayerError> {
Ok(self
.claim_stale_page(stream, group, consumer, start_id, config)
.await?
.entries)
}
async fn claim_stale_page(
&self,
stream: &str,
group: &str,
consumer: &str,
start_id: &str,
config: RuntimeQueueReclaimConfig,
) -> Result<RuntimeQueueReclaimPage, DataLayerError> {
validate_runtime_queue_name(stream, "runtime queue stream")?;
validate_runtime_queue_name(group, "runtime queue group")?;
validate_runtime_queue_name(consumer, "runtime queue consumer")?;
@@ -1250,32 +1384,76 @@ impl RuntimeQueueStore for RuntimeState {
match self.backend.as_ref() {
RuntimeStateBackend::Memory(memory) => {
memory
.queue_claim_stale(stream, group, consumer, start_id, config)
.queue_claim_stale_page(stream, group, consumer, start_id, config)
.await
}
RuntimeStateBackend::Redis(redis) => Ok(redis
.stream
.claim_stale(
&RedisStreamName(stream.to_string()),
&RedisConsumerGroup(group.to_string()),
&RedisConsumerName(consumer.to_string()),
start_id,
RedisStreamReclaimConfig {
min_idle_ms: config.min_idle_ms,
count: config.count,
},
)
.await?
.entries
.into_iter()
.map(|entry| RuntimeQueueEntry {
id: entry.id,
fields: entry.fields,
RuntimeStateBackend::Redis(redis) => {
let page = redis
.stream
.claim_stale(
&RedisStreamName(stream.to_string()),
&RedisConsumerGroup(group.to_string()),
&RedisConsumerName(consumer.to_string()),
start_id,
RedisStreamReclaimConfig {
min_idle_ms: config.min_idle_ms,
count: config.count,
},
)
.await?;
Ok(RuntimeQueueReclaimPage {
next_start_id: page.next_start_id,
entries: page
.entries
.into_iter()
.map(|entry| RuntimeQueueEntry {
id: entry.id,
fields: entry.fields,
})
.collect(),
deleted_ids: page.deleted_ids,
})
.collect()),
}
}
}
async fn try_transfer_pending_to_stream(
&self,
source: &str,
group: &str,
entry_id: &str,
destination: &str,
destination_fields: &BTreeMap<String, String>,
) -> Result<Option<RuntimeQueueTransferOutcome>, DataLayerError> {
validate_runtime_queue_transfer(source, group, entry_id, destination, destination_fields)?;
let outcome = match self.backend.as_ref() {
RuntimeStateBackend::Memory(memory) => {
memory
.queue_transfer_pending_to_stream(
source,
group,
entry_id,
destination,
destination_fields,
)
.await?
}
RuntimeStateBackend::Redis(redis) => {
redis
.stream
.try_transfer_pending_to_stream(
source,
group,
entry_id,
destination,
destination_fields,
)
.await?
}
};
Ok(Some(outcome))
}
async fn ack(
&self,
stream: &str,
@@ -1817,6 +1995,18 @@ mod tests {
use std::process::{Child, Command, Stdio};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
mod stream_receive {
include!("redis/stream_receive_tests.rs");
}
mod dead_letter_transfer {
include!("redis/dead_letter_transfer_tests.rs");
}
mod usage_limit_cleanup {
include!("redis/usage_limit_cleanup_tests.rs");
}
#[tokio::test]
async fn memory_kv_expires_entries() {
let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default());
@@ -2722,6 +2912,131 @@ mod tests {
}
}
#[tokio::test]
async fn redis_large_stream_batches_preserve_fields_across_read_reclaim_and_ack() {
let Some(redis) = TestRedisServer::start().await else {
return;
};
for protocol in ["resp2", "resp3"] {
let runtime = RuntimeState::redis(
RedisClientConfig {
url: format!("{}?protocol={protocol}", redis.redis_url),
key_prefix: Some(format!("large-batch-{protocol}")),
},
Some(5_000),
)
.await
.expect("large batch runtime should connect");
let stream = "usage:large-batch";
let group = "workers";
RuntimeQueueStore::ensure_consumer_group(&runtime, stream, group, "0-0")
.await
.unwrap();
let payload = format!(
"{}\r\n\"escaped\"\\\u{4e2d}\u{6587}",
"x".repeat(512 * 1024)
);
let mut expected = BTreeMap::new();
for sequence in 0..24 {
let fields = BTreeMap::from([
("payload".to_string(), payload.clone()),
("sequence".to_string(), sequence.to_string()),
("legacy_marker".to_string(), "preserve exactly".to_string()),
]);
let id =
RuntimeQueueStore::append_fields_with_maxlen(&runtime, stream, &fields, None)
.await
.unwrap();
expected.insert(id, sequence.to_string());
}
let mut readers = tokio::task::JoinSet::new();
for index in 0..3 {
let runtime = runtime.clone();
readers.spawn(async move {
RuntimeQueueStore::read_group(
&runtime,
stream,
group,
&format!("reader-{index}"),
8,
Some(1),
)
.await
.unwrap()
});
}
let mut delivered = std::collections::BTreeSet::new();
while let Some(entries) = readers.join_next().await {
let entries = entries.unwrap();
assert_eq!(entries.len(), 8);
for entry in entries {
assert_eq!(entry.fields.len(), 3);
assert_eq!(entry.fields["payload"].as_bytes(), payload.as_bytes());
assert_eq!(entry.fields["sequence"], expected[&entry.id]);
assert_eq!(entry.fields["legacy_marker"], "preserve exactly");
assert!(delivered.insert(entry.id));
}
}
assert_eq!(delivered.len(), 24);
let stats = RuntimeQueueStore::stats(&runtime, stream, Some(group))
.await
.unwrap();
assert_eq!(stats.group_pending, 24);
assert_eq!(stats.group_lag, Some(0));
tokio::time::sleep(Duration::from_millis(20)).await;
let mut reclaimed = std::collections::BTreeSet::new();
while reclaimed.len() < 24 {
let entries = RuntimeQueueStore::claim_stale(
&runtime,
stream,
group,
"retry-consumer",
"0-0",
RuntimeQueueReclaimConfig {
min_idle_ms: 1,
count: 5,
},
)
.await
.unwrap();
assert!(!entries.is_empty());
assert!(entries.len() <= 5);
let mut ids = Vec::new();
for entry in entries {
assert_eq!(entry.fields.len(), 3);
assert_eq!(entry.fields["payload"].as_bytes(), payload.as_bytes());
assert_eq!(entry.fields["sequence"], expected[&entry.id]);
assert_eq!(entry.fields["legacy_marker"], "preserve exactly");
assert!(reclaimed.insert(entry.id.clone()));
ids.push(entry.id);
}
assert_eq!(
RuntimeQueueStore::ack(&runtime, stream, group, &ids)
.await
.unwrap(),
ids.len()
);
assert_eq!(
RuntimeQueueStore::delete(&runtime, stream, &ids)
.await
.unwrap(),
ids.len()
);
}
assert_eq!(reclaimed, delivered);
let stats = RuntimeQueueStore::stats(&runtime, stream, Some(group))
.await
.unwrap();
assert_eq!(stats.stream_length, 0);
assert_eq!(stats.group_pending, 0);
assert_eq!(stats.group_lag, Some(0));
eprintln!(
"verified {protocol}: 24 large records, 3 readers, read/reclaim/ack complete"
);
}
}
#[tokio::test]
async fn redis_connection_manager_recovers_after_restart() {
let Some(mut redis) = TestRedisServer::start().await else {
@@ -2776,6 +3091,269 @@ mod tests {
assert_kv_score_and_queue_contract(&redis_runtime).await;
}
#[tokio::test]
async fn runtime_backends_share_bounded_score_window_aggregation() {
let memory = RuntimeState::memory(MemoryRuntimeStateConfig::default());
assert_bounded_score_window_aggregation(&memory).await;
let Some((_server, runtime)) = redis_runtime_for_test("score-window").await else {
return;
};
assert_bounded_score_window_aggregation(&runtime).await;
}
async fn assert_bounded_score_window_aggregation(runtime: &RuntimeState) {
let keys = (0..35)
.map(|index| format!("window:{index}"))
.collect::<Vec<_>>();
for (member, score) in [
("expired:999", 99.999),
("boundary:7", 100.0),
("recent:9007199254740993", 101.0),
("nested:prefix:+00012", 102.0),
("zero:0", 103.0),
("invalid:1.5", 104.0),
("invalid:-1", 105.0),
("invalid:18446744073709551616", 106.0),
("invalid: 12", 107.0),
("missing-separator", 108.0),
] {
runtime
.score_set(&keys[0], member, score)
.await
.expect("seed values");
}
runtime
.score_set(&keys[1], "max:18446744073709551615", 100.0)
.await
.expect("seed max");
runtime
.score_set(&keys[2], "max:18446744073709551615", 100.0)
.await
.expect("seed overflow");
runtime
.score_set(&keys[2], "additional:2", 100.0)
.await
.expect("seed overflow addition");
for index in 0..SCORE_WINDOW_AGGREGATION_MEMBER_LIMIT {
runtime
.score_set(&keys[3], &format!("{index}:3"), 100.0)
.await
.expect("seed bounded window");
}
runtime
.score_set(&keys[3], "expired:9999", 0.0)
.await
.expect("seed expired sample");
let stats = runtime
.score_window_u64_stats_by_min(&keys, 100.0)
.await
.expect("aggregate");
assert_eq!(
stats.len(),
keys.len(),
"pipeline batches preserve key order"
);
assert_eq!(
stats[0],
Some(ScoreWindowU64Stats {
sum: 9_007_199_254_741_012,
positive_count: 3
})
);
assert_eq!(
stats[1],
Some(ScoreWindowU64Stats {
sum: u64::MAX,
positive_count: 1
})
);
assert_eq!(
stats[2],
Some(ScoreWindowU64Stats {
sum: u64::MAX,
positive_count: 2
})
);
assert_eq!(
stats[3],
Some(ScoreWindowU64Stats {
sum: 1536,
positive_count: 512
})
);
assert!(stats[4..]
.iter()
.all(|stats| *stats == Some(ScoreWindowU64Stats::default())));
runtime
.score_set(&keys[3], "overflowing-window:11", 101.0)
.await
.expect("exceed server limit");
let stats = runtime
.score_window_u64_stats_by_min(&keys[3..4], 100.0)
.await
.expect("bounded fallback");
assert_eq!(
stats,
vec![None],
"oversized windows require the full exact read"
);
let members = runtime
.score_range_by_min(&keys[3], 100.0)
.await
.expect("full window");
assert_eq!(
ScoreWindowU64Stats::from_members(members.iter().map(String::as_str)).sum,
1547
);
runtime
.score_remove(&keys[3], "overflowing-window:11")
.await
.expect("remove newest");
assert_eq!(
runtime
.score_window_u64_stats_by_min(&keys[3..4], 100.0)
.await
.expect("read after remove")[0]
.unwrap()
.sum,
1536
);
runtime
.score_set(&keys[3], "0:3", 99.0)
.await
.expect("move sample outside window");
assert_eq!(
runtime
.score_window_u64_stats_by_min(&keys[3..4], 100.0)
.await
.expect("read changed score")[0]
.unwrap()
.sum,
1533
);
assert!(runtime
.score_window_u64_stats_by_min(&[], 100.0)
.await
.expect("empty query")
.is_empty());
assert!(runtime
.score_window_u64_stats_by_min(&keys, f64::NAN)
.await
.is_err());
}
#[tokio::test]
async fn redis_score_window_aggregation_reloads_scripts_without_caching_old_cost() {
let Some((server, runtime)) = redis_runtime_for_test("score-window-reload").await else {
return;
};
let keys = vec!["reload:cost".to_string()];
runtime
.score_set(&keys[0], "first:7", 100.0)
.await
.expect("first cost");
assert_eq!(
runtime
.score_window_u64_stats_by_min(&keys, 100.0)
.await
.expect("first aggregate")[0]
.unwrap()
.sum,
7
);
let client = ::redis::Client::open(server.redis_url.as_str()).expect("test Redis client");
let mut connection = client
.get_multiplexed_async_connection()
.await
.expect("test connection");
::redis::cmd("SCRIPT")
.arg("FLUSH")
.query_async::<()>(&mut connection)
.await
.expect("flush scripts");
runtime
.score_set(&keys[0], "second:11", 101.0)
.await
.expect("new cost");
assert_eq!(
runtime
.score_window_u64_stats_by_min(&keys, 100.0)
.await
.expect("reload aggregate")[0]
.unwrap()
.sum,
18
);
assert_eq!(
runtime
.score_window_u64_stats_by_min(&keys, 101.0)
.await
.expect("changed window")[0]
.unwrap()
.sum,
11
);
runtime
.key_expire(&keys[0], Duration::ZERO)
.await
.expect("expire window");
assert_eq!(
runtime
.score_window_u64_stats_by_min(&keys, 100.0)
.await
.expect("expired aggregate"),
vec![Some(ScoreWindowU64Stats::default())]
);
}
#[tokio::test]
async fn redis_score_window_aggregation_observes_completed_concurrent_writes() {
let Some((_server, runtime)) = redis_runtime_for_test("score-window-concurrent").await
else {
return;
};
let writer_runtime = runtime.clone();
let (written_tx, mut written_rx) = tokio::sync::mpsc::channel(8);
let writer = tokio::spawn(async move {
for index in 1..=128_u64 {
writer_runtime
.score_set("concurrent:cost", &format!("{index}:2"), 100.0)
.await
.expect("concurrent write");
written_tx.send(index).await.expect("notify reader");
}
});
let keys = vec!["concurrent:cost".to_string()];
while let Some(written) = written_rx.recv().await {
let stats = runtime
.score_window_u64_stats_by_min(&keys, 100.0)
.await
.expect("concurrent aggregate")[0]
.unwrap();
assert!(
stats.positive_count >= written,
"completed writes must not be hidden by a stale aggregate"
);
assert_eq!(
stats.sum,
stats.positive_count * 2,
"one script observes one consistent window"
);
}
writer.await.expect("writer task");
assert_eq!(
runtime
.score_window_u64_stats_by_min(&keys, 100.0)
.await
.expect("final aggregate")[0]
.unwrap()
.sum,
256
);
}
#[tokio::test]
async fn runtime_backends_reject_invalid_shared_inputs() {
let memory = RuntimeState::memory(MemoryRuntimeStateConfig::default());
File diff suppressed because it is too large Load Diff
+255 -21
View File
@@ -1,9 +1,12 @@
use crate::error::RedisResultExt;
use crate::redis::RedisKeyspace;
use crate::DataLayerError;
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use std::sync::{Arc, Mutex};
use std::time::Duration;
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
use tracing::info;
pub(crate) type RedisClient = redis::Client;
@@ -141,8 +144,8 @@ pub(crate) struct RedisConnectionRouter {
fast: RedisManagedConnection,
stream: Arc<Vec<RedisManagedConnection>>,
stream_next: Arc<AtomicUsize>,
blocking_stream: Arc<Vec<RedisManagedConnection>>,
blocking_stream_next: Arc<AtomicUsize>,
blocking_stream: Arc<RedisBlockingStreamPool>,
usage_cleanup: Arc<RedisBlockingStreamPool>,
admin: RedisManagedConnection,
metrics: Arc<RedisConnectionMetrics>,
}
@@ -152,7 +155,7 @@ impl std::fmt::Debug for RedisConnectionRouter {
f.debug_struct("RedisConnectionRouter")
.field("lanes", &["fast", "stream", "blocking_stream", "admin"])
.field("stream_lanes", &self.stream.len())
.field("blocking_stream_lanes", &self.blocking_stream.len())
.field("blocking_stream_lanes", &self.blocking_stream.capacity)
.finish()
}
}
@@ -182,7 +185,7 @@ impl RedisConnectionRouter {
)
.await?;
let stream_lanes = stream.len();
let blocking_stream_lanes = blocking_stream.len();
let blocking_stream_lanes = blocking_stream.capacity;
info!(
redis_lanes = "fast,stream,blocking_stream,admin",
redis_stream_lanes = stream_lanes,
@@ -194,7 +197,11 @@ impl RedisConnectionRouter {
stream: Arc::new(stream),
stream_next: Arc::new(AtomicUsize::new(0)),
blocking_stream: Arc::new(blocking_stream),
blocking_stream_next: Arc::new(AtomicUsize::new(0)),
usage_cleanup: Arc::new(RedisBlockingStreamPool::new(
client,
command_timeout_ms,
vec![None, None],
)),
admin,
metrics: Arc::new(RedisConnectionMetrics::default()),
})
@@ -208,13 +215,26 @@ impl RedisConnectionRouter {
self.stream[index].clone()
}
RedisConnectionLane::BlockingStream => {
let index = next_lane_index(&self.blocking_stream_next, self.blocking_stream.len());
self.blocking_stream[index].clone()
unreachable!("blocking stream commands require an exclusive connection lease")
}
RedisConnectionLane::Admin => self.admin.clone(),
}
}
pub(crate) async fn blocking_stream_connection(
&self,
) -> Result<RedisBlockingStreamLease, DataLayerError> {
self.blocking_stream.checkout().await
}
// WATCH/MULTI state must never share a multiplexed connection with other callers.
// These lazy leases own their drivers, so cancellation closes an unfinished transaction.
pub(crate) async fn usage_cleanup_connection(
&self,
) -> Result<RedisBlockingStreamLease, DataLayerError> {
self.usage_cleanup.checkout().await
}
pub(crate) fn record_error(&self, lane: RedisConnectionLane) {
self.metrics
.for_lane(lane)
@@ -257,6 +277,114 @@ impl RedisConnectionRouter {
}
}
struct RedisBlockingConnection {
connection: redis::aio::MultiplexedConnection,
driver: Pin<Box<dyn Future<Output = ()> + Send>>,
}
impl RedisBlockingConnection {
async fn query(&mut self, command: &redis::Cmd) -> Result<redis::Value, DataLayerError> {
// Drive the connection inside its owning query, so cancellation drops the
// socket immediately instead of leaving a spawned driver with an old BLOCK.
tokio::select! {
result = command.query_async(&mut self.connection) => result.map_redis_err(),
() = self.driver.as_mut() => Err(DataLayerError::Redis(
"runtime redis blocking stream connection driver terminated".to_string(),
)),
}
}
}
struct RedisBlockingStreamPool {
client: RedisClient,
command_timeout_ms: Option<u64>,
capacity: usize,
available: Mutex<Vec<Option<RedisBlockingConnection>>>,
permits: Arc<Semaphore>,
}
impl RedisBlockingStreamPool {
fn new(
client: RedisClient,
command_timeout_ms: Option<u64>,
available: Vec<Option<RedisBlockingConnection>>,
) -> Self {
let capacity = available.len();
Self {
client,
command_timeout_ms,
capacity,
available: Mutex::new(available),
permits: Arc::new(Semaphore::new(capacity)),
}
}
async fn checkout(self: &Arc<Self>) -> Result<RedisBlockingStreamLease, DataLayerError> {
let permit = Arc::clone(&self.permits)
.acquire_owned()
.await
.map_err(|_| {
DataLayerError::Redis("runtime redis blocking stream pool closed".to_string())
})?;
let connection = self
.available
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.pop()
.expect("blocking stream permit must have an available slot");
Ok(RedisBlockingStreamLease {
pool: Arc::clone(self),
connection,
reusable: false,
_permit: permit,
})
}
}
pub(crate) struct RedisBlockingStreamLease {
pool: Arc<RedisBlockingStreamPool>,
connection: Option<RedisBlockingConnection>,
reusable: bool,
_permit: OwnedSemaphorePermit,
}
impl RedisBlockingStreamLease {
pub(crate) async fn query(
&mut self,
command: &redis::Cmd,
) -> Result<redis::Value, DataLayerError> {
self.reusable = false;
if self.connection.is_none() {
self.connection = Some(
connect_blocking_stream_lane(&self.pool.client, self.pool.command_timeout_ms)
.await?,
);
}
self.connection
.as_mut()
.expect("blocking stream connection initialized")
.query(command)
.await
}
pub(crate) fn recycle(&mut self) {
self.reusable = true;
}
}
impl Drop for RedisBlockingStreamLease {
fn drop(&mut self) {
let connection = self.connection.take().filter(|_| self.reusable);
self.pool
.available
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.push(connection);
// The permit is released after the slot is restored. An uncompleted
// query has already dropped both its connection and its owned driver.
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize)]
pub struct RedisLaneDiagnostics {
pub lane: &'static str,
@@ -406,21 +534,46 @@ async fn connect_blocking_stream_lanes(
client: &RedisClient,
command_timeout_ms: Option<u64>,
requested_lanes: Option<usize>,
) -> Result<Vec<RedisManagedConnection>, DataLayerError> {
) -> Result<RedisBlockingStreamPool, DataLayerError> {
let lane_count = blocking_stream_lane_count(requested_lanes)?;
let mut lanes = Vec::with_capacity(lane_count);
for _ in 0..lane_count {
lanes.push(
connect_lane(
client,
connection_manager_config(command_timeout_ms),
RedisConnectionLane::BlockingStream,
command_timeout_ms,
)
.await?,
);
lanes.push(Some(
connect_blocking_stream_lane(client, command_timeout_ms).await?,
));
}
Ok(lanes)
Ok(RedisBlockingStreamPool::new(
client.clone(),
command_timeout_ms,
lanes,
))
}
async fn connect_blocking_stream_lane(
client: &RedisClient,
command_timeout_ms: Option<u64>,
) -> Result<RedisBlockingConnection, DataLayerError> {
let connect = client.create_multiplexed_tokio_connection();
let result = if let Some(timeout_ms) = command_timeout_ms {
tokio::time::timeout(Duration::from_millis(timeout_ms), connect)
.await
.map_err(|_| {
DataLayerError::TimedOut(format!(
"runtime redis blocking_stream lane connection exceeded {timeout_ms}ms timeout"
))
})?
} else {
connect.await
};
let (connection, driver) = result.map_err(|err| {
DataLayerError::Redis(format!(
"failed to initialize runtime redis blocking_stream lane: {err}"
))
})?;
Ok(RedisBlockingConnection {
connection,
driver: Box::pin(driver),
})
}
fn blocking_stream_lane_count(requested_lanes: Option<usize>) -> Result<usize, DataLayerError> {
@@ -480,10 +633,12 @@ async fn connect_lane(
mod tests {
use super::{
blocking_stream_lane_count, default_blocking_stream_lane_count, next_lane_index,
stream_lane_count, RedisClientConfig, RedisClientFactory, RedisLaneMetrics,
DEFAULT_STREAM_LANES, MAX_BLOCKING_STREAM_LANES_CAP, REDIS_COMMAND_LATENCY_BUCKETS_MS,
stream_lane_count, RedisBlockingStreamPool, RedisClientConfig, RedisClientFactory,
RedisLaneMetrics, DEFAULT_STREAM_LANES, MAX_BLOCKING_STREAM_LANES_CAP,
REDIS_COMMAND_LATENCY_BUCKETS_MS,
};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
#[test]
@@ -543,6 +698,85 @@ mod tests {
assert!(blocking_stream_lane_count(Some(0)).is_err());
}
fn empty_blocking_pool(capacity: usize) -> Arc<RedisBlockingStreamPool> {
Arc::new(RedisBlockingStreamPool::new(
redis::Client::open("redis://127.0.0.1/0").expect("lazy client"),
None,
(0..capacity).map(|_| None).collect(),
))
}
#[tokio::test]
async fn blocking_stream_pool_cancelled_checkout_preserves_owner_and_capacity() {
use std::future::Future;
use std::task::Poll;
let pool = empty_blocking_pool(1);
let owner = pool.checkout().await.expect("first lease");
let mut waiting = Box::pin(pool.checkout());
std::future::poll_fn(|cx| {
assert!(matches!(waiting.as_mut().poll(cx), Poll::Pending));
Poll::Ready(())
})
.await;
drop(waiting);
assert_eq!(pool.permits.available_permits(), 0);
assert!(pool.available.lock().unwrap().is_empty());
drop(owner);
assert_eq!(pool.permits.available_permits(), 1);
let replacement = pool
.checkout()
.await
.expect("cancelled owner slot restored");
assert!(replacement.connection.is_none());
let panic = tokio::spawn(async move {
let _lease = replacement;
panic!("test owner panic");
});
assert!(panic.await.expect_err("owner panicked").is_panic());
assert_eq!(pool.permits.available_permits(), 1);
assert_eq!(pool.available.lock().unwrap().len(), 1);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn blocking_stream_pool_concurrent_checkouts_stay_within_capacity() {
let pool = empty_blocking_pool(3);
let active = Arc::new(AtomicUsize::new(0));
let peak = Arc::new(AtomicUsize::new(0));
let barrier = Arc::new(tokio::sync::Barrier::new(32));
let mut tasks = tokio::task::JoinSet::new();
for _ in 0..32 {
let pool = Arc::clone(&pool);
let active = Arc::clone(&active);
let peak = Arc::clone(&peak);
let barrier = Arc::clone(&barrier);
tasks.spawn(async move {
barrier.wait().await;
for _ in 0..32 {
let lease = pool.checkout().await.expect("bounded lease");
let concurrent = active.fetch_add(1, Ordering::SeqCst) + 1;
peak.fetch_max(concurrent, Ordering::SeqCst);
assert!(concurrent <= 3);
tokio::task::yield_now().await;
active.fetch_sub(1, Ordering::SeqCst);
drop(lease);
}
});
}
tokio::time::timeout(Duration::from_secs(5), async {
while let Some(result) = tasks.join_next().await {
result.expect("checkout task");
}
})
.await
.expect("all checkouts complete without losing capacity");
assert!(peak.load(Ordering::SeqCst) <= 3);
assert_eq!(active.load(Ordering::SeqCst), 0);
assert_eq!(pool.permits.available_permits(), 3);
assert_eq!(pool.available.lock().unwrap().len(), 3);
}
#[test]
fn stream_lane_count_uses_fixed_default() {
assert_eq!(stream_lane_count(), DEFAULT_STREAM_LANES);
@@ -0,0 +1,44 @@
-- Redis scripts do not roll back errors. Complete all predictable checks before XADD.
if #KEYS ~= 2 or #ARGV < 4 or (#ARGV - 2) % 2 ~= 0 then
return redis.error_reply('ERR invalid pending transfer arguments')
end
if KEYS[1] == KEYS[2] then
return redis.error_reply('ERR pending transfer source and destination must differ')
end
if type(redis.acl_check_cmd) ~= 'function' then
return redis.error_reply('ERR pending transfer requires Redis 7 or later for ACL preflight')
end
-- The exact range bounds keep work independent of the size of the PEL.
local pending = redis.call('XPENDING', KEYS[1], ARGV[1], ARGV[2], ARGV[2], 1)
if #pending == 0 then
return {0, '', 0, 0}
end
if pending[1][1] ~= ARGV[2] then
return redis.error_reply('ERR pending transfer requires an exact canonical entry ID')
end
local destination_type = redis.call('TYPE', KEYS[2]).ok
if destination_type ~= 'none' and destination_type ~= 'stream' then
return redis.error_reply('WRONGTYPE pending transfer destination must be a stream')
end
local append = {'XADD', KEYS[2], '*'}
for index = 3, #ARGV do
append[#append + 1] = ARGV[index]
end
if not redis.acl_check_cmd(unpack(append)) then
return redis.error_reply('NOPERM pending transfer requires XADD permission')
end
if not redis.acl_check_cmd('XACK', KEYS[1], ARGV[1], ARGV[2]) then
return redis.error_reply('NOPERM pending transfer requires XACK permission')
end
if not redis.acl_check_cmd('XDEL', KEYS[1], ARGV[2]) then
return redis.error_reply('NOPERM pending transfer requires XDEL permission')
end
-- PEL membership, stream types and ACLs cannot change between these commands.
-- A trimmed source body is still recoverable from the caller's retained fields.
local destination_id = redis.call(unpack(append))
local acked = redis.call('XACK', KEYS[1], ARGV[1], ARGV[2])
local deleted = redis.call('XDEL', KEYS[1], ARGV[2])
return {1, destination_id, acked, deleted}
@@ -0,0 +1,500 @@
use super::*;
type TransferTestConnection = ::redis::aio::MultiplexedConnection;
const TRANSFER_GROUP: &str = "transfer-workers";
const TRANSFER_USER: &str = "transfer-worker";
async fn transfer_runtime(
protocol: &str,
) -> Option<(TestRedisServer, RuntimeState, TransferTestConnection)> {
let Some(server) = TestRedisServer::start().await else {
eprintln!(
"dead letter transfer {protocol} skipped: isolated Redis fixture unavailable; check AETHER_REDIS_SERVER_BIN"
);
return None;
};
let mut admin = ::redis::Client::open(server.redis_url.clone())
.expect("transfer admin client")
.get_multiplexed_async_connection()
.await
.expect("transfer admin connection");
::redis::cmd("ACL")
.arg("SETUSER")
.arg(TRANSFER_USER)
.arg("on")
.arg(">transfer-test-password")
.arg("~*")
.arg("+@all")
.query_async::<()>(&mut admin)
.await
.expect("transfer test user");
let runtime = RuntimeState::redis_with_blocking_stream_lanes(
RedisClientConfig {
url: format!(
"redis://{TRANSFER_USER}:[email protected]:{}/5?protocol={protocol}",
server.port
),
key_prefix: Some(format!("transfer-{protocol}")),
},
Some(5_000),
Some(4),
)
.await
.expect("authenticated transfer runtime");
::redis::cmd("SELECT")
.arg(5)
.query_async::<()>(&mut admin)
.await
.expect("transfer admin database");
eprintln!(
"dead letter transfer fixture ready: protocol={protocol} db=5 port={} authenticated=true",
server.port
);
Some((server, runtime, admin))
}
fn transfer_source_fields() -> BTreeMap<String, String> {
BTreeMap::from([
(
"payload".to_string(),
"malformed\r\n\"quoted\"\\\u{4e2d}\u{6587}\0".to_string(),
),
("legacy".to_string(), "retain every field".to_string()),
(
String::new(),
"empty field name is valid Redis data".to_string(),
),
])
}
fn transfer_archive_fields(entry: &RuntimeQueueEntry) -> BTreeMap<String, String> {
BTreeMap::from([
(
"payload".to_string(),
serde_json::json!({
"entry_id": entry.id,
"fields": entry.fields,
"error": "invalid record\r\n\"details\"\\\u{4e2d}\u{6587}"
})
.to_string(),
),
("archive_version".to_string(), "1".to_string()),
])
}
async fn seed_transfer_entry(runtime: &RuntimeState, source: &str) -> RuntimeQueueEntry {
RuntimeQueueStore::ensure_consumer_group(runtime, source, TRANSFER_GROUP, "0-0")
.await
.expect("source consumer group");
let id = RuntimeQueueStore::append_fields_with_maxlen(
runtime,
source,
&transfer_source_fields(),
None,
)
.await
.expect("source append");
let mut entries =
RuntimeQueueStore::read_group(runtime, source, TRANSFER_GROUP, "owner", 1, None)
.await
.expect("pending source entry");
assert_eq!(entries.len(), 1);
let entry = entries.pop().expect("one entry");
assert_eq!(entry.id, id);
assert_eq!(entry.fields, transfer_source_fields());
entry
}
async fn transfer_entries(
admin: &mut TransferTestConnection,
stream: &str,
) -> Vec<RuntimeQueueEntry> {
let rows = ::redis::cmd("XRANGE")
.arg(stream)
.arg("-")
.arg("+")
.query_async::<::redis::streams::StreamRangeReply>(admin)
.await
.expect("inspect transfer stream");
rows.ids
.into_iter()
.map(|row| RuntimeQueueEntry {
id: row.id,
fields: row
.map
.into_iter()
.map(|(field, value)| {
let value =
::redis::from_redis_value::<String>(&value).expect("string field value");
(field, value)
})
.collect(),
})
.collect()
}
async fn transfer_pending(
admin: &mut TransferTestConnection,
source: &str,
) -> Vec<(String, String, u64, u64)> {
::redis::cmd("XPENDING")
.arg(source)
.arg(TRANSFER_GROUP)
.arg("-")
.arg("+")
.arg(16)
.query_async(admin)
.await
.expect("inspect transfer pending entries")
}
async fn transfer_entry(
runtime: &RuntimeState,
source: &str,
entry: &RuntimeQueueEntry,
destination: &str,
fields: &BTreeMap<String, String>,
) -> Result<RuntimeQueueTransferOutcome, DataLayerError> {
RuntimeQueueStore::try_transfer_pending_to_stream(
runtime,
source,
TRANSFER_GROUP,
&entry.id,
destination,
fields,
)
.await
.map(|outcome| outcome.expect("Redis implements atomic transfer"))
}
async fn assert_transfer_source_unchanged(
admin: &mut TransferTestConnection,
source: &str,
entry: &RuntimeQueueEntry,
) {
assert_eq!(transfer_entries(admin, source).await, [entry.clone()]);
let pending = transfer_pending(admin, source).await;
assert_eq!(pending.len(), 1);
assert_eq!(pending[0].0, entry.id);
assert_eq!(pending[0].1, "owner");
}
async fn assert_transfer_completed(
admin: &mut TransferTestConnection,
source: &str,
destination: &str,
fields: &BTreeMap<String, String>,
) {
assert!(transfer_entries(admin, source).await.is_empty());
assert!(transfer_pending(admin, source).await.is_empty());
let archived = transfer_entries(admin, destination).await;
assert_eq!(archived.len(), 1);
assert_eq!(&archived[0].fields, fields);
}
#[tokio::test]
async fn redis_dead_letter_transfer_concurrent_consumers_archive_exactly_once() {
for protocol in ["resp2", "resp3"] {
let Some((_server, runtime, mut admin)) = transfer_runtime(protocol).await else {
return;
};
let source = "usage:{transfer}:concurrent";
let destination = "usage:{transfer}:concurrent:dlq";
let entry = seed_transfer_entry(&runtime, source).await;
let fields = transfer_archive_fields(&entry);
let barrier = Arc::new(tokio::sync::Barrier::new(16));
let mut tasks = tokio::task::JoinSet::new();
for _ in 0..16 {
let runtime = runtime.clone();
let entry = entry.clone();
let fields = fields.clone();
let barrier = Arc::clone(&barrier);
tasks.spawn(async move {
barrier.wait().await;
transfer_entry(&runtime, source, &entry, destination, &fields)
.await
.expect("concurrent transfer")
});
}
let mut transferred_ids = Vec::new();
let mut not_pending = 0;
while let Some(result) = tasks.join_next().await {
match result.expect("transfer task") {
RuntimeQueueTransferOutcome::Transferred {
destination_id,
acked,
deleted,
} => {
assert_eq!((acked, deleted), (1, 1));
transferred_ids.push(destination_id);
}
RuntimeQueueTransferOutcome::NotPending => not_pending += 1,
}
}
assert_eq!(not_pending, 15);
assert_eq!(transferred_ids.len(), 1);
assert_transfer_completed(&mut admin, source, destination, &fields).await;
assert_eq!(
transfer_entries(&mut admin, destination).await[0].id,
transferred_ids[0]
);
let diagnostics = runtime.redis_diagnostics().await.unwrap().unwrap();
let lane = diagnostics
.lanes
.iter()
.find(|lane| lane.lane == "blocking_stream")
.expect("exclusive stream lane diagnostics");
assert!(lane.command_count >= 16);
}
}
#[tokio::test]
async fn redis_dead_letter_transfer_retry_after_ignored_success_does_not_archive_again() {
let Some((_server, runtime, mut admin)) = transfer_runtime("resp2").await else {
return;
};
let source = "usage:{transfer}:retry";
let destination = "usage:{transfer}:retry:dlq";
let entry = seed_transfer_entry(&runtime, source).await;
let fields = transfer_archive_fields(&entry);
// Commit the operation, but discard its result as a caller missing the reply would.
let _ = transfer_entry(&runtime, source, &entry, destination, &fields)
.await
.expect("first transfer commits");
let changed_fields = BTreeMap::from([("payload".to_string(), "retry value".to_string())]);
assert_eq!(
transfer_entry(&runtime, source, &entry, destination, &changed_fields)
.await
.expect("retry is successful"),
RuntimeQueueTransferOutcome::NotPending
);
assert_transfer_completed(&mut admin, source, destination, &fields).await;
}
#[tokio::test]
async fn redis_dead_letter_transfer_acl_preflight_prevents_partial_writes() {
let Some((_server, runtime, mut admin)) = transfer_runtime("resp3").await else {
return;
};
for forbidden in ["XADD", "XACK", "XDEL"] {
let source = format!("usage:{{transfer}}:acl-{forbidden}");
let destination = format!("{source}:dlq");
let entry = seed_transfer_entry(&runtime, &source).await;
let fields = transfer_archive_fields(&entry);
::redis::cmd("ACL")
.arg("SETUSER")
.arg(TRANSFER_USER)
.arg(format!("-{forbidden}"))
.query_async::<()>(&mut admin)
.await
.expect("deny one write command");
let error = transfer_entry(&runtime, &source, &entry, &destination, &fields)
.await
.expect_err("denied write must fail before archiving");
assert!(
error
.to_string()
.contains(&format!("requires {forbidden} permission")),
"expected the {forbidden} preflight error, got {error}"
);
assert_transfer_source_unchanged(&mut admin, &source, &entry).await;
assert!(transfer_entries(&mut admin, &destination).await.is_empty());
::redis::cmd("ACL")
.arg("SETUSER")
.arg(TRANSFER_USER)
.arg(format!("+{forbidden}"))
.query_async::<()>(&mut admin)
.await
.expect("restore one write command");
assert!(matches!(
transfer_entry(&runtime, &source, &entry, &destination, &fields)
.await
.expect("retry after restoring permission"),
RuntimeQueueTransferOutcome::Transferred {
acked: 1,
deleted: 1,
..
}
));
assert_eq!(
transfer_entry(&runtime, &source, &entry, &destination, &fields)
.await
.expect("idempotent retry"),
RuntimeQueueTransferOutcome::NotPending
);
assert_transfer_completed(&mut admin, &source, &destination, &fields).await;
}
}
#[tokio::test]
async fn redis_dead_letter_transfer_invalid_state_preserves_source_until_repaired() {
let Some((_server, runtime, mut admin)) = transfer_runtime("resp2").await else {
return;
};
let source = "usage:{transfer}:invalid";
let destination = "usage:{transfer}:invalid:dlq";
let entry = seed_transfer_entry(&runtime, source).await;
let fields = transfer_archive_fields(&entry);
::redis::cmd("SET")
.arg(destination)
.arg("existing non-stream data")
.query_async::<()>(&mut admin)
.await
.expect("wrong-type destination");
let error = transfer_entry(&runtime, source, &entry, destination, &fields)
.await
.expect_err("wrong type must not acknowledge source");
assert!(error.to_string().contains("WRONGTYPE"));
assert_transfer_source_unchanged(&mut admin, source, &entry).await;
assert_eq!(
::redis::cmd("GET")
.arg(destination)
.query_async::<String>(&mut admin)
.await
.unwrap(),
"existing non-stream data"
);
::redis::cmd("DEL")
.arg(destination)
.query_async::<usize>(&mut admin)
.await
.expect("repair destination type");
let error = RuntimeQueueStore::try_transfer_pending_to_stream(
&runtime,
source,
"missing-group",
&entry.id,
destination,
&fields,
)
.await
.expect_err("missing group must fail before archiving");
assert!(error.to_string().contains("NOGROUP"));
for invalid_id in [
"",
"-",
"+",
"1",
"01-0",
"1-00",
"(1-0",
"1-+0",
"18446744073709551616-0",
] {
assert!(matches!(
RuntimeQueueStore::try_transfer_pending_to_stream(
&runtime,
source,
TRANSFER_GROUP,
invalid_id,
destination,
&fields,
)
.await,
Err(DataLayerError::InvalidInput(_))
));
}
assert!(matches!(
transfer_entry(&runtime, source, &entry, source, &fields).await,
Err(DataLayerError::InvalidInput(_))
));
assert!(matches!(
transfer_entry(&runtime, source, &entry, destination, &BTreeMap::new()).await,
Err(DataLayerError::InvalidInput(_))
));
assert_transfer_source_unchanged(&mut admin, source, &entry).await;
assert!(transfer_entries(&mut admin, destination).await.is_empty());
assert!(matches!(
transfer_entry(&runtime, source, &entry, destination, &fields)
.await
.expect("valid transfer after failed attempts"),
RuntimeQueueTransferOutcome::Transferred {
acked: 1,
deleted: 1,
..
}
));
assert_transfer_completed(&mut admin, source, destination, &fields).await;
}
#[tokio::test]
async fn redis_dead_letter_transfer_archives_retained_fields_after_source_trim() {
for protocol in ["resp2", "resp3"] {
let Some((_server, runtime, mut admin)) = transfer_runtime(protocol).await else {
return;
};
let source = "usage:{transfer}:trimmed";
let destination = "usage:{transfer}:trimmed:dlq";
let entry = seed_transfer_entry(&runtime, source).await;
let fields = transfer_archive_fields(&entry);
let trimmed = ::redis::cmd("XTRIM")
.arg(source)
.arg("MAXLEN")
.arg(0)
.query_async::<usize>(&mut admin)
.await
.expect("trim source body while retaining PEL");
assert_eq!(trimmed, 1);
assert!(transfer_entries(&mut admin, source).await.is_empty());
assert_eq!(transfer_pending(&mut admin, source).await[0].0, entry.id);
assert!(matches!(
transfer_entry(&runtime, source, &entry, destination, &fields)
.await
.expect("pending body remains recoverable from caller fields"),
RuntimeQueueTransferOutcome::Transferred {
acked: 1,
deleted: 0,
..
}
));
assert_eq!(
transfer_entry(&runtime, source, &entry, destination, &fields)
.await
.expect("trimmed entry retry"),
RuntimeQueueTransferOutcome::NotPending
);
assert_transfer_completed(&mut admin, source, destination, &fields).await;
}
}
#[tokio::test]
async fn redis_dead_letter_transfer_preserves_ids_larger_than_lua_integer_precision() {
let Some((_server, runtime, mut admin)) = transfer_runtime("resp3").await else {
return;
};
let source = "usage:{transfer}:large-id";
let destination = "usage:{transfer}:large-id:dlq";
let entry_id = "9007199254740993-18446744073709551614";
RuntimeQueueStore::ensure_consumer_group(&runtime, source, TRANSFER_GROUP, "0-0")
.await
.expect("large-ID source group");
::redis::cmd("XADD")
.arg(source)
.arg(entry_id)
.arg("payload")
.arg("retained value")
.query_async::<String>(&mut admin)
.await
.expect("large stream ID");
let mut entries =
RuntimeQueueStore::read_group(&runtime, source, TRANSFER_GROUP, "owner", 1, None)
.await
.expect("read large-ID entry");
assert_eq!(entries.len(), 1);
let entry = entries.pop().unwrap();
assert_eq!(entry.id, entry_id);
let fields = transfer_archive_fields(&entry);
assert!(matches!(
transfer_entry(&runtime, source, &entry, destination, &fields)
.await
.expect("transfer exact large ID"),
RuntimeQueueTransferOutcome::Transferred {
acked: 1,
deleted: 1,
..
}
));
assert_transfer_completed(&mut admin, source, destination, &fields).await;
}
@@ -4,6 +4,7 @@ mod lock;
mod namespace;
mod runtime;
mod stream;
mod usage_cleanup;
pub use client::{RedisClientConfig, RedisLaneDiagnostics};
pub use kv::{RedisKvRunner, RedisKvRunnerConfig};
+186 -25
View File
@@ -7,9 +7,13 @@ use crate::redis::{
};
use crate::{
DataLayerError, RateLimitCheck, RateLimitInput, RateLimitScope, RuntimeSemaphoreError,
UsageLimitCheck, UsageLimitInput, UsageLimitReleaseInput,
ScoreWindowU64Stats, UsageLimitCheck, UsageLimitInput, UsageLimitReleaseInput,
SCORE_WINDOW_AGGREGATION_MEMBER_LIMIT,
};
const SCORE_WINDOW_STATS_SCRIPT: &str = include_str!("score_window.lua");
const SCORE_WINDOW_STATS_PIPELINE_KEY_LIMIT: usize = 16;
const RATE_LIMIT_CHECK_AND_CONSUME_SCRIPT: &str = r#"
local user_key = KEYS[1]
local key_key = KEYS[2]
@@ -52,15 +56,119 @@ end
return {1, 0, 0, remaining}
"#;
const USAGE_LIMIT_CHECK_AND_CONSUME_SCRIPT: &str = r#"
pub(super) const USAGE_LIMIT_CHECK_AND_CONSUME_SCRIPT: &str = r#"
local count = #KEYS
local now = tonumber(ARGV[1])
local event_id = ARGV[2]
-- Large mixed windows are copied in bounded commands on an exclusive WATCH
-- connection. This read-only pass must finish before pruning any of the rules.
if ARGV[count * 3 + 3] ~= 'inline' and redis.acl_check_cmd then
local deferred = {2}
for i = 1, count do
local key = KEYS[i]
if redis.call('ZCARD', key) > 4096 then
local cutoff = now - tonumber(ARGV[(i - 1) * 3 + 4]) * 1000
local expired = redis.pcall('ZCOUNT', key, '-inf', cutoff)
if type(expired) == 'number' and expired > 4096 then
local live = redis.call('ZCARD', key) - expired
local temporary = key .. ':__usage_copy:acl'
if live > 256
and redis.acl_check_cmd('WATCH', key)
and redis.acl_check_cmd('UNWATCH')
and redis.acl_check_cmd('MULTI')
and redis.acl_check_cmd('EXEC')
and redis.acl_check_cmd('EVAL', 'return 1', 0)
and redis.acl_check_cmd('EXISTS', temporary)
and redis.acl_check_cmd('ZRANGE', key, 0, 511, 'WITHSCORES')
and redis.acl_check_cmd('ZADD', temporary, 0, 'acl')
and redis.acl_check_cmd('PTTL', key)
and redis.acl_check_cmd('PEXPIRE', temporary, 60000)
and redis.acl_check_cmd('PERSIST', temporary)
and redis.acl_check_cmd('UNLINK', key, temporary)
and redis.acl_check_cmd('RENAME', temporary, key) then
deferred[#deferred + 1] = i
deferred[#deferred + 1] = expired
deferred[#deferred + 1] = live
end
end
end
end
if #deferred > 1 then return deferred end
end
local function replace_with_survivors(key, live)
if not redis.acl_check_cmd then return false end
local temporary = key .. ':__usage_trim'
local exists = redis.pcall('EXISTS', temporary)
local ttl = redis.pcall('PTTL', key)
if exists ~= 0 or type(ttl) ~= 'number' or ttl == 0 or ttl < -1 then
return false
end
local rows = redis.call('ZRANGE', key, -live, -1, 'WITHSCORES')
local args = {}
for i = 1, #rows, 2 do
args[#args + 1] = rows[i + 1]
args[#args + 1] = rows[i]
end
-- Validate every write before detaching the original. The temporary key
-- is fully built first; restricted ACLs keep the original cleanup path.
if not redis.acl_check_cmd('ZADD', temporary, unpack(args))
or not redis.acl_check_cmd('UNLINK', key)
or not redis.acl_check_cmd('UNLINK', temporary)
or not redis.acl_check_cmd('RENAME', temporary, key)
or (ttl > 0 and not redis.acl_check_cmd('PEXPIRE', temporary, ttl)) then
return false
end
if type(redis.pcall('ZADD', temporary, unpack(args))) ~= 'number' then
return false
end
if ttl > 0 and redis.pcall('PEXPIRE', temporary, ttl) ~= 1 then
redis.call('UNLINK', temporary)
return false
end
if type(redis.pcall('UNLINK', key)) ~= 'number' then
redis.call('UNLINK', temporary)
return false
end
redis.call('RENAME', temporary, key)
return true
end
local function prune_window(key, cutoff)
local cardinality = redis.call('ZCARD', key)
if cardinality == 0 then return 0 end
if cardinality > 256 then
local earliest = redis.call('ZRANGE', key, 0, 0, 'WITHSCORES')
if #earliest >= 2 and tonumber(earliest[2]) > cutoff then
return cardinality
end
local latest = redis.call('ZRANGE', key, -1, -1, 'WITHSCORES')
if #latest >= 2 and tonumber(latest[2]) <= cutoff then
-- The entire window is expired. Redis can free the detached object
-- off its command thread while this key is reused.
local result = redis.pcall('UNLINK', key)
if type(result) == 'number' then return 0 end
elseif cardinality > 1024 then
local expired = redis.pcall('ZCOUNT', key, '-inf', cutoff)
if type(expired) == 'number' then
local live = cardinality - expired
if live > 0 and live <= 256 and expired > live * 4
and replace_with_survivors(key, live) then
return live
end
end
end
end
-- Preserve the existing path for large live windows and restricted ACLs.
-- Partial deferred deletion could resurrect entries on out-of-order calls.
return cardinality - redis.call('ZREMRANGEBYSCORE', key, '-inf', cutoff)
end
local current_counts = {}
for i = 1, count do
local window_ms = tonumber(ARGV[(i - 1) * 3 + 4]) * 1000
local cutoff = now - window_ms
redis.call('ZREMRANGEBYSCORE', KEYS[i], '-inf', cutoff)
current_counts[i] = prune_window(KEYS[i], now - window_ms)
end
for i = 1, count do
@@ -68,7 +176,7 @@ for i = 1, count do
local window_ms = tonumber(ARGV[(i - 1) * 3 + 4]) * 1000
local already_consumed = redis.call('ZSCORE', KEYS[i], event_id)
if not already_consumed then
local current = redis.call('ZCARD', KEYS[i])
local current = current_counts[i]
if current >= limit then
local earliest = redis.call('ZRANGE', KEYS[i], 0, 0, 'WITHSCORES')
local retry_after = 1
@@ -362,30 +470,17 @@ impl RedisRuntimeRunner {
&self,
input: UsageLimitInput<'_>,
) -> Result<UsageLimitCheck, DataLayerError> {
let script = script(USAGE_LIMIT_CHECK_AND_CONSUME_SCRIPT);
let mut invocation = script.prepare_invoke();
for rule in input.rules {
invocation.key(self.keyspace.key(rule.key));
}
invocation.arg(input.now_unix_ms as i64);
invocation.arg(input.event_id);
for rule in input.rules {
invocation.arg(rule.limit as i64);
invocation.arg(rule.window_seconds as i64);
invocation.arg(rule.retention_seconds as i64);
}
let keys = input
.rules
.iter()
.map(|rule| self.keyspace.key(rule.key))
.collect::<Vec<_>>();
let raw = run_lane_with_timeout(
&self.connections,
RedisConnectionLane::Fast,
self.command_timeout_ms,
Some(self.command_timeout_ms.unwrap_or(30_000)),
"runtime usage limit check",
async {
let mut connection = self.connections.connection(RedisConnectionLane::Fast);
invocation
.invoke_async::<Vec<i64>>(&mut connection)
.await
.map_redis_err()
},
super::usage_cleanup::check_and_consume(&self.connections, &keys, &input),
)
.await?;
match raw.first().copied() {
@@ -539,6 +634,72 @@ impl RedisRuntimeRunner {
.await
}
pub(crate) async fn score_window_u64_stats_by_min(
&self,
keys: &[String],
min_score: f64,
) -> Result<Vec<Option<ScoreWindowU64Stats>>, DataLayerError> {
let script = script(SCORE_WINDOW_STATS_SCRIPT);
let mut output = Vec::with_capacity(keys.len());
for batch in keys.chunks(SCORE_WINDOW_STATS_PIPELINE_KEY_LIMIT) {
let values: Vec<(u8, String, u64)> = run_lane_with_timeout(
&self.connections,
RedisConnectionLane::Admin,
self.command_timeout_ms,
"runtime score window stats",
async {
let mut connection = self.connections.connection(RedisConnectionLane::Admin);
let mut pipeline = redis::pipe();
for key in batch {
pipeline
.cmd("EVALSHA")
.arg(script.get_hash())
.arg(1)
.arg(self.keyspace.key(key))
.arg(min_score)
.arg(SCORE_WINDOW_AGGREGATION_MEMBER_LIMIT);
}
match pipeline.query_async(&mut connection).await {
Err(err) if err.kind() == redis::ErrorKind::NoScriptError => {
script
.prepare_invoke()
.load_async(&mut connection)
.await
.map_redis_err()?;
pipeline.query_async(&mut connection).await.map_redis_err()
}
result => result.map_redis_err(),
}
},
)
.await?;
if values.len() != batch.len() {
return Err(DataLayerError::UnexpectedValue(
"runtime score window stats result count mismatch".to_string(),
));
}
for (aggregated, sum, positive_count) in values {
output.push(match aggregated {
0 => None,
1 => Some(ScoreWindowU64Stats {
sum: sum.parse().map_err(|_| {
DataLayerError::UnexpectedValue(
"runtime score window stats returned an invalid sum".to_string(),
)
})?,
positive_count,
}),
_ => {
return Err(DataLayerError::UnexpectedValue(
"runtime score window stats returned an invalid status".to_string(),
))
}
});
}
}
Ok(output)
}
pub(crate) async fn score_remove_by_score(
&self,
key: &str,
@@ -0,0 +1,36 @@
-- Count before reading members so a large window cannot run an unbounded Lua loop.
local count = redis.call('ZCOUNT', KEYS[1], ARGV[1], '+inf')
if count > tonumber(ARGV[2]) then
return {0, '0', 0}
end
local members = redis.call('ZRANGEBYSCORE', KEYS[1], ARGV[1], '+inf')
local high, low, positive = 0, 0, 0
local max_high, max_low = 18446744073, 709551615
for _, member in ipairs(members) do
local value = string.match(member, ':([^:]*)$')
if value and string.match(value, '^%+?%d+$') then
value = string.gsub(value, '^%+', '')
value = string.gsub(value, '^0+', '')
if #value > 0 and (#value < 20 or (#value == 20 and value <= '18446744073709551615')) then
positive = positive + 1
-- Two base-1e9 limbs keep every integer operation exactly representable
-- in Redis Lua's doubles, including values above 2^53 and u64::MAX.
local split = math.max(0, #value - 9)
local value_high = tonumber(string.sub(value, 1, split)) or 0
local value_low = tonumber(string.sub(value, split + 1))
low = low + value_low
high = high + value_high + math.floor(low / 1000000000)
low = low % 1000000000
if high > max_high or (high == max_high and low > max_low) then
high, low = max_high, max_low
end
end
end
end
local total = string.format('%.0f', low)
if high > 0 then
total = string.format('%.0f', high) .. string.format('%09d', low)
end
return {1, total, positive}
+190 -50
View File
@@ -1,16 +1,19 @@
use std::collections::BTreeMap;
use std::collections::{BTreeMap, HashMap};
use std::future::Future;
use redis::from_redis_value;
use redis::streams::StreamReadReply;
use redis::Value as RedisValue;
use redis::{from_owned_redis_value, from_redis_value};
use crate::error::{redis_error, RedisResultExt};
use crate::redis::{
run_lane_with_timeout, RedisClientConfig, RedisClientFactory, RedisConnectionLane,
RedisConnectionRouter, RedisKeyspace,
};
use crate::{DataLayerError, RuntimeQueueStats};
use crate::{
validate_runtime_queue_transfer, DataLayerError, RuntimeQueueStats, RuntimeQueueTransferOutcome,
};
const DEAD_LETTER_TRANSFER_SCRIPT: &str = include_str!("dead_letter_transfer.lua");
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct RedisStreamName(pub String);
@@ -259,6 +262,41 @@ impl RedisStreamRunner {
self.append_fields(stream, &fields).await
}
/// Atomically archives a pending entry and removes it from its source group and stream.
/// Requires Redis 7+ for ACL preflight. Both stream keys must share a slot on Redis Cluster.
pub async fn try_transfer_pending_to_stream(
&self,
source: &str,
group: &str,
entry_id: &str,
destination: &str,
destination_fields: &BTreeMap<String, String>,
) -> Result<RuntimeQueueTransferOutcome, DataLayerError> {
validate_runtime_queue_transfer(source, group, entry_id, destination, destination_fields)?;
self.run_with_timeout(
RedisConnectionLane::BlockingStream,
"redis stream pending transfer",
async {
let mut lease = self.connections.blocking_stream_connection().await?;
let mut command = redis::cmd("EVAL");
command
.arg(DEAD_LETTER_TRANSFER_SCRIPT)
.arg(2)
.arg(source)
.arg(destination)
.arg(group)
.arg(entry_id);
for (field, value) in destination_fields {
command.arg(field).arg(value);
}
let outcome = parse_transfer_result(lease.query(&command).await?)?;
lease.recycle();
Ok(outcome)
},
)
.await
}
pub async fn read_group(
&self,
stream: &RedisStreamName,
@@ -275,7 +313,6 @@ impl RedisStreamRunner {
RedisConnectionLane::Stream
};
self.run_with_timeout(lane, "redis stream read group", async {
let mut connection = self.connections.connection(lane);
let mut command = redis::cmd("XREADGROUP");
command
.arg("GROUP")
@@ -288,6 +325,14 @@ impl RedisStreamRunner {
}
command.arg("STREAMS").arg(&stream.0).arg(">");
if lane == RedisConnectionLane::BlockingStream {
let mut lease = self.connections.blocking_stream_connection().await?;
let entries = parse_stream_read_entries(lease.query(&command).await?)?;
lease.recycle();
return Ok(entries);
}
let mut connection = self.connections.connection(lane);
let reply = command
.query_async::<RedisValue>(&mut connection)
.await
@@ -364,22 +409,25 @@ impl RedisStreamRunner {
validate_stream_position(start_id)?;
config.validate()?;
self.run_with_timeout(RedisConnectionLane::Stream, "redis stream reclaim", async {
let mut connection = self.connections.connection(RedisConnectionLane::Stream);
let reply = redis::cmd("XAUTOCLAIM")
.arg(&stream.0)
.arg(&group.0)
.arg(&consumer.0)
.arg(config.min_idle_ms)
.arg(start_id)
.arg("COUNT")
.arg(config.count)
.query_async::<RedisValue>(&mut connection)
.await
.map_redis_err()?;
parse_reclaim_result(reply)
})
self.run_with_timeout(
RedisConnectionLane::BlockingStream,
"redis stream reclaim",
async {
let mut lease = self.connections.blocking_stream_connection().await?;
let mut command = redis::cmd("XAUTOCLAIM");
command
.arg(&stream.0)
.arg(&group.0)
.arg(&consumer.0)
.arg(config.min_idle_ms)
.arg(start_id)
.arg("COUNT")
.arg(config.count);
let result = parse_reclaim_result(lease.query(&command).await?)?;
lease.recycle();
Ok(result)
},
)
.await
}
@@ -536,18 +584,21 @@ fn parse_stream_read_entries(value: RedisValue) -> Result<Vec<RedisStreamEntry>,
return Ok(Vec::new());
}
let reply = from_redis_value::<StreamReadReply>(&value).map_err(redis_error)?;
Ok(reply
.keys
// StreamReadReply in redis 0.28 falls back to borrowed conversion even for an
// owned input. Its underlying containers support moving every payload buffer.
type StreamReadRows = Vec<HashMap<String, Vec<HashMap<String, HashMap<String, RedisValue>>>>>;
let rows = from_owned_redis_value::<StreamReadRows>(value).map_err(redis_error)?;
Ok(rows
.into_iter()
.flat_map(|key| key.ids.into_iter())
.map(|id| RedisStreamEntry {
id: id.id,
fields: id
.map
.flat_map(HashMap::into_values)
.flatten()
.flat_map(HashMap::into_iter)
.map(|(id, fields)| RedisStreamEntry {
id,
fields: fields
.into_iter()
.filter_map(|(field, value)| {
redis::from_redis_value::<String>(&value)
from_owned_redis_value::<String>(value)
.ok()
.map(|text| (field, text))
})
@@ -556,6 +607,31 @@ fn parse_stream_read_entries(value: RedisValue) -> Result<Vec<RedisStreamEntry>,
.collect())
}
fn parse_transfer_result(value: RedisValue) -> Result<RuntimeQueueTransferOutcome, DataLayerError> {
if !matches!(&value, RedisValue::Array(parts) if parts.len() == 4) {
return Err(DataLayerError::UnexpectedValue(
"redis stream pending transfer returned an invalid result shape".to_string(),
));
}
let (transferred, destination_id, acked, deleted) =
from_owned_redis_value::<(i64, String, usize, usize)>(value).map_err(redis_error)?;
match transferred {
0 if destination_id.is_empty() && acked == 0 && deleted == 0 => {
Ok(RuntimeQueueTransferOutcome::NotPending)
}
1 if !destination_id.is_empty() && acked == 1 && deleted <= 1 => {
Ok(RuntimeQueueTransferOutcome::Transferred {
destination_id,
acked,
deleted,
})
}
_ => Err(DataLayerError::UnexpectedValue(
"redis stream pending transfer returned inconsistent outcome fields".to_string(),
)),
}
}
fn parse_reclaim_result(value: RedisValue) -> Result<RedisStreamReclaimResult, DataLayerError> {
let RedisValue::Array(parts) = value else {
return Err(DataLayerError::UnexpectedValue(
@@ -570,9 +646,13 @@ fn parse_reclaim_result(value: RedisValue) -> Result<RedisStreamReclaimResult, D
)));
}
let next_start_id = parse_string_value(&parts[0], "redis xautoclaim next_start_id")?;
let entries = parse_reclaim_entries(&parts[1])?;
let deleted_ids = match parts.get(2) {
let mut parts = parts.into_iter();
let next_start_id = parse_owned_string_value(
parts.next().expect("validated reclaim result length"),
"redis xautoclaim next_start_id",
)?;
let entries = parse_reclaim_entries(parts.next().expect("validated reclaim result length"))?;
let deleted_ids = match parts.next() {
Some(value) => parse_string_array(value, "redis xautoclaim deleted_ids")?,
None => Vec::new(),
};
@@ -584,9 +664,9 @@ fn parse_reclaim_result(value: RedisValue) -> Result<RedisStreamReclaimResult, D
})
}
fn parse_reclaim_entries(value: &RedisValue) -> Result<Vec<RedisStreamEntry>, DataLayerError> {
fn parse_reclaim_entries(value: RedisValue) -> Result<Vec<RedisStreamEntry>, DataLayerError> {
match value {
RedisValue::Array(entries) => entries.iter().map(parse_reclaim_entry).collect(),
RedisValue::Array(entries) => entries.into_iter().map(parse_reclaim_entry).collect(),
RedisValue::Nil => Ok(Vec::new()),
_ => Err(DataLayerError::UnexpectedValue(
"redis xautoclaim entries payload was not an array".to_string(),
@@ -594,7 +674,7 @@ fn parse_reclaim_entries(value: &RedisValue) -> Result<Vec<RedisStreamEntry>, Da
}
}
fn parse_reclaim_entry(value: &RedisValue) -> Result<RedisStreamEntry, DataLayerError> {
fn parse_reclaim_entry(value: RedisValue) -> Result<RedisStreamEntry, DataLayerError> {
let RedisValue::Array(parts) = value else {
return Err(DataLayerError::UnexpectedValue(
"redis xautoclaim entry was not an array".to_string(),
@@ -607,13 +687,20 @@ fn parse_reclaim_entry(value: &RedisValue) -> Result<RedisStreamEntry, DataLayer
)));
}
let id = parse_string_value(&parts[0], "redis xautoclaim entry id")?;
let fields = parse_string_map(&parts[1], "redis xautoclaim entry fields")?;
let mut parts = parts.into_iter();
let id = parse_owned_string_value(
parts.next().expect("validated reclaim entry length"),
"redis xautoclaim entry id",
)?;
let fields = parse_string_map(
parts.next().expect("validated reclaim entry length"),
"redis xautoclaim entry fields",
)?;
Ok(RedisStreamEntry { id, fields })
}
fn parse_string_map(
value: &RedisValue,
value: RedisValue,
context: &str,
) -> Result<BTreeMap<String, String>, DataLayerError> {
match value {
@@ -625,19 +712,23 @@ fn parse_string_map(
)));
}
let mut fields = BTreeMap::new();
for pair in values.chunks(2) {
let key = parse_string_value(&pair[0], context)?;
let value = parse_string_value(&pair[1], context)?;
let mut values = values.into_iter();
while let Some(key) = values.next() {
let key = parse_owned_string_value(key, context)?;
let value = parse_owned_string_value(
values.next().expect("validated even number of fields"),
context,
)?;
fields.insert(key, value);
}
Ok(fields)
}
RedisValue::Map(entries) => entries
.iter()
.into_iter()
.map(|(key, value)| {
Ok((
parse_string_value(key, context)?,
parse_string_value(value, context)?,
parse_owned_string_value(key, context)?,
parse_owned_string_value(value, context)?,
))
})
.collect(),
@@ -648,11 +739,11 @@ fn parse_string_map(
}
}
fn parse_string_array(value: &RedisValue, context: &str) -> Result<Vec<String>, DataLayerError> {
fn parse_string_array(value: RedisValue, context: &str) -> Result<Vec<String>, DataLayerError> {
match value {
RedisValue::Array(values) => values
.iter()
.map(|value| parse_string_value(value, context))
.into_iter()
.map(|value| parse_owned_string_value(value, context))
.collect(),
RedisValue::Nil => Ok(Vec::new()),
_ => Err(DataLayerError::UnexpectedValue(format!(
@@ -669,6 +760,14 @@ fn parse_string_value(value: &RedisValue, context: &str) -> Result<String, DataL
})
}
fn parse_owned_string_value(value: RedisValue, context: &str) -> Result<String, DataLayerError> {
from_owned_redis_value::<String>(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"{context} was not a string-compatible redis value: {err}"
))
})
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
struct RedisXInfoGroupStats {
pending: Option<u64>,
@@ -788,19 +887,60 @@ fn parse_u64_value(value: &RedisValue, context: &str) -> Result<u64, DataLayerEr
})
}
#[cfg(test)]
#[path = "stream_owned_reply_tests.rs"]
mod owned_reply_tests;
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use super::{
parse_reclaim_result, parse_stream_read_entries, parse_xinfo_group_stats,
parse_xpending_oldest_idle_ms, redis_stream_stats_missing_stream,
parse_reclaim_result, parse_stream_read_entries, parse_transfer_result,
parse_xinfo_group_stats, parse_xpending_oldest_idle_ms, redis_stream_stats_missing_stream,
redis_stream_stats_missing_stream_or_group, validate_consumer, validate_group,
validate_stream_name, validate_stream_position, RedisConsumerName, RedisStreamName,
RedisStreamReclaimConfig, RedisStreamReclaimResult, RedisStreamRunnerConfig,
};
use redis::Value as RedisValue;
#[test]
fn pending_transfer_result_requires_complete_consistent_outcome() {
assert!(matches!(
parse_transfer_result(RedisValue::Array(vec![
RedisValue::Int(0),
RedisValue::BulkString(Vec::new()),
RedisValue::Int(0),
RedisValue::Int(0),
])),
Ok(crate::RuntimeQueueTransferOutcome::NotPending)
));
for values in [
Vec::new(),
vec![RedisValue::Int(0)],
vec![
RedisValue::Int(0),
RedisValue::BulkString(b"1-0".to_vec()),
RedisValue::Int(0),
RedisValue::Int(0),
],
vec![
RedisValue::Int(1),
RedisValue::BulkString(b"1-0".to_vec()),
RedisValue::Int(0),
RedisValue::Int(1),
],
vec![
RedisValue::Int(1),
RedisValue::BulkString(b"1-0".to_vec()),
RedisValue::Int(1),
RedisValue::Int(2),
],
] {
assert!(parse_transfer_result(RedisValue::Array(values)).is_err());
}
}
#[test]
fn validates_stream_runner_config() {
assert!(RedisStreamRunnerConfig {
@@ -0,0 +1,410 @@
use std::collections::BTreeMap;
use redis::streams::StreamReadReply;
use redis::{from_redis_value, Value as RedisValue, VerbatimFormat};
use super::{parse_reclaim_result, parse_stream_read_entries, RedisStreamEntry};
use crate::DataLayerError;
fn bulk(text: &str) -> RedisValue {
RedisValue::BulkString(text.as_bytes().to_vec())
}
fn fields_reply(fields: Vec<(RedisValue, RedisValue)>, resp3: bool) -> RedisValue {
if resp3 {
RedisValue::Map(fields)
} else {
RedisValue::Array(
fields
.into_iter()
.flat_map(|(key, value)| [key, value])
.collect(),
)
}
}
fn read_reply(id: RedisValue, fields: RedisValue, resp3: bool) -> RedisValue {
let entries = RedisValue::Array(vec![RedisValue::Array(vec![id, fields])]);
if resp3 {
RedisValue::Map(vec![(bulk("usage:events"), entries)])
} else {
RedisValue::Array(vec![RedisValue::Array(vec![bulk("usage:events"), entries])])
}
}
fn reclaim_reply(id: RedisValue, fields: RedisValue) -> RedisValue {
RedisValue::Array(vec![
bulk("0-0"),
RedisValue::Array(vec![RedisValue::Array(vec![id, fields])]),
RedisValue::Nil,
])
}
fn original_read_parser(value: &RedisValue) -> Result<Vec<RedisStreamEntry>, DataLayerError> {
let reply = from_redis_value::<StreamReadReply>(value).map_err(crate::error::redis_error)?;
Ok(reply
.keys
.into_iter()
.flat_map(|key| key.ids)
.map(|entry| RedisStreamEntry {
id: entry.id,
fields: entry
.map
.into_iter()
.filter_map(|(field, value)| {
from_redis_value::<String>(&value)
.ok()
.map(|value| (field, value))
})
.collect(),
})
.collect())
}
#[test]
fn owned_read_moves_large_payload_and_id_buffers_in_resp2_and_resp3() {
for resp3 in [false, true] {
let mut payload = Vec::with_capacity(512 * 1024);
payload.extend_from_slice(b"{\"message\":\"");
payload.resize(256 * 1024, b'x');
payload.extend_from_slice(b"\",\"cache_read_input_tokens\":0}\n");
let expected = payload.clone();
let pointer = payload.as_ptr();
let capacity = payload.capacity();
let mut id = Vec::with_capacity(64);
id.extend_from_slice(b"1710000000000-0");
let id_pointer = id.as_ptr();
let id_capacity = id.capacity();
let fields = fields_reply(
vec![(bulk("payload"), RedisValue::BulkString(payload))],
resp3,
);
let parsed =
parse_stream_read_entries(read_reply(RedisValue::BulkString(id), fields, resp3))
.expect("read reply");
assert_eq!(parsed.len(), 1);
let body = &parsed[0].fields["payload"];
assert_eq!(body.as_bytes(), expected);
assert_eq!(
body.as_ptr(),
pointer,
"the RESP payload allocation must be reused"
);
assert_eq!(body.capacity(), capacity);
assert_eq!(parsed[0].id.as_ptr(), id_pointer);
assert_eq!(parsed[0].id.capacity(), id_capacity);
}
}
#[test]
fn owned_reclaim_moves_payload_and_all_id_buffers() {
for resp3 in [false, true] {
let mut payload = Vec::with_capacity(512 * 1024);
payload.resize(256 * 1024, b'x');
let pointer = payload.as_ptr();
let capacity = payload.capacity();
let mut id = Vec::with_capacity(64);
id.extend_from_slice(b"1710000000000-0");
let id_pointer = id.as_ptr();
let mut next_id = Vec::with_capacity(64);
next_id.extend_from_slice(b"1710000000001-0");
let next_pointer = next_id.as_ptr();
let mut deleted_id = Vec::with_capacity(64);
deleted_id.extend_from_slice(b"1709999999999-0");
let deleted_pointer = deleted_id.as_ptr();
let fields = fields_reply(
vec![(bulk("payload"), RedisValue::BulkString(payload))],
resp3,
);
let parsed = parse_reclaim_result(RedisValue::Array(vec![
RedisValue::BulkString(next_id),
RedisValue::Array(vec![RedisValue::Array(vec![
RedisValue::BulkString(id),
fields,
])]),
RedisValue::Array(vec![RedisValue::BulkString(deleted_id)]),
]))
.expect("reclaim reply");
let body = &parsed.entries[0].fields["payload"];
assert_eq!(body.len(), 256 * 1024);
assert!(body.bytes().all(|byte| byte == b'x'));
assert_eq!(body.as_ptr(), pointer);
assert_eq!(body.capacity(), capacity);
assert_eq!(parsed.entries[0].id.as_ptr(), id_pointer);
assert_eq!(parsed.next_start_id.as_ptr(), next_pointer);
assert_eq!(parsed.deleted_ids[0].as_ptr(), deleted_pointer);
}
}
#[test]
fn owned_read_matches_existing_redis_decoder_for_supported_reply_shapes() {
let mut replies = vec![
RedisValue::Nil,
RedisValue::Array(vec![]),
RedisValue::Map(vec![]),
];
for resp3 in [false, true] {
for fields in [
RedisValue::Nil,
fields_reply(
vec![(
bulk("payload"),
bulk("{ \"text\": \"caf\u{00e9}\", \"n\": 0 }\n"),
)],
resp3,
),
fields_reply(
vec![
(bulk("payload"), RedisValue::Nil),
(bulk("count"), RedisValue::Int(0)),
],
resp3,
),
fields_reply(
vec![(bulk("payload"), RedisValue::BulkString(vec![0xff]))],
resp3,
),
fields_reply(
vec![
(bulk("payload"), bulk("first")),
(bulk("payload"), bulk("last")),
(bulk("invalid"), RedisValue::Boolean(false)),
],
resp3,
),
fields_reply(
vec![(
bulk("payload"),
RedisValue::Attribute {
data: Box::new(bulk("annotated payload")),
attributes: vec![(bulk("encoding"), bulk("utf8"))],
},
)],
resp3,
),
] {
replies.push(read_reply(bulk("1-0"), fields, resp3));
}
replies.push(read_reply(RedisValue::Int(42), RedisValue::Nil, resp3));
}
for reply in replies {
let expected = original_read_parser(&reply).expect("baseline reply");
assert_eq!(
parse_stream_read_entries(reply).expect("owned reply"),
expected
);
}
}
#[test]
fn owned_read_preserves_duplicate_overwrite_before_value_filtering() {
for resp3 in [false, true] {
let reply = read_reply(
bulk("1-0"),
fields_reply(
vec![
(bulk("payload"), bulk("valid earlier payload")),
(bulk("payload"), RedisValue::BulkString(vec![0xff])),
(bulk("retry"), bulk("first")),
(bulk("retry"), bulk("last")),
],
resp3,
),
resp3,
);
let parsed = parse_stream_read_entries(reply)
.expect("invalid values are filtered after deduplication");
assert_eq!(
parsed[0].fields,
BTreeMap::from([("retry".to_string(), "last".to_string())])
);
}
}
#[test]
fn owned_read_keeps_invalid_id_key_and_shape_errors_in_redis_category() {
for reply in [
RedisValue::Int(7),
RedisValue::Array(vec![RedisValue::Array(vec![bulk("stream")])]),
read_reply(RedisValue::BulkString(vec![0xff]), RedisValue::Nil, false),
read_reply(bulk("1-0"), RedisValue::Array(vec![bulk("orphan")]), false),
read_reply(
bulk("1-0"),
fields_reply(
vec![(RedisValue::BulkString(vec![0xff]), bulk("value"))],
true,
),
true,
),
] {
assert!(matches!(
original_read_parser(&reply),
Err(DataLayerError::Redis(_))
));
assert!(matches!(
parse_stream_read_entries(reply),
Err(DataLayerError::Redis(_))
));
}
}
#[test]
fn owned_reclaim_preserves_string_types_nil_and_duplicate_fields() {
for resp3 in [false, true] {
let fields = fields_reply(
vec![
(bulk("payload"), bulk("first")),
(bulk("payload"), bulk("{\"text\":\"caf\u{00e9}\"}\n")),
(bulk("zero"), RedisValue::Int(0)),
(bulk("double"), RedisValue::Double(1.5)),
(
bulk("simple"),
RedisValue::SimpleString("simple".to_string()),
),
(bulk("okay"), RedisValue::Okay),
(
bulk("verbatim"),
RedisValue::VerbatimString {
format: VerbatimFormat::Text,
text: "verbatim".to_string(),
},
),
(
bulk("attribute"),
RedisValue::Attribute {
data: Box::new(bulk("annotated")),
attributes: vec![],
},
),
],
resp3,
);
let parsed =
parse_reclaim_result(reclaim_reply(bulk("1-0"), fields)).expect("reclaim reply");
assert_eq!(
parsed.entries[0].fields,
BTreeMap::from([
(
"payload".to_string(),
"{\"text\":\"caf\u{00e9}\"}\n".to_string()
),
("zero".to_string(), "0".to_string()),
("double".to_string(), "1.5".to_string()),
("simple".to_string(), "simple".to_string()),
("okay".to_string(), "OK".to_string()),
("verbatim".to_string(), "verbatim".to_string()),
("attribute".to_string(), "annotated".to_string()),
])
);
assert!(parsed.deleted_ids.is_empty());
}
let parsed =
parse_reclaim_result(reclaim_reply(bulk("1-0"), RedisValue::Nil)).expect("nil fields");
assert!(parsed.entries[0].fields.is_empty());
let parsed = parse_reclaim_result(RedisValue::Array(vec![bulk("0-0"), RedisValue::Nil]))
.expect("nil entries");
assert!(parsed.entries.is_empty());
assert!(parsed.deleted_ids.is_empty());
}
#[test]
fn owned_reclaim_preserves_strict_validation_and_error_context() {
for (reply, context) in [
(
RedisValue::Nil,
"redis xautoclaim returned non-array payload",
),
(
RedisValue::Array(vec![bulk("0-0")]),
"redis xautoclaim returned 1 top-level fields",
),
(
RedisValue::Array(vec![RedisValue::Nil, RedisValue::Nil]),
"redis xautoclaim next_start_id",
),
(
RedisValue::Array(vec![bulk("0-0"), RedisValue::Int(1)]),
"redis xautoclaim entries payload was not an array",
),
(
RedisValue::Array(vec![bulk("0-0"), RedisValue::Array(vec![RedisValue::Nil])]),
"redis xautoclaim entry was not an array",
),
(
RedisValue::Array(vec![
bulk("0-0"),
RedisValue::Array(vec![RedisValue::Array(vec![])]),
]),
"redis xautoclaim entry had 0 fields",
),
(
reclaim_reply(RedisValue::BulkString(vec![0xff]), RedisValue::Nil),
"redis xautoclaim entry id",
),
(
reclaim_reply(bulk("1-0"), RedisValue::Array(vec![bulk("orphan")])),
"redis xautoclaim entry fields expected an even number",
),
(
reclaim_reply(bulk("1-0"), RedisValue::Int(1)),
"redis xautoclaim entry fields expected a redis array/map payload",
),
(
reclaim_reply(
bulk("1-0"),
fields_reply(
vec![(bulk("payload"), RedisValue::BulkString(vec![0xff]))],
false,
),
),
"redis xautoclaim entry fields was not a string-compatible",
),
(
reclaim_reply(
bulk("1-0"),
fields_reply(
vec![(RedisValue::BulkString(vec![0xff]), bulk("value"))],
true,
),
),
"redis xautoclaim entry fields was not a string-compatible",
),
(
RedisValue::Array(vec![bulk("0-0"), RedisValue::Nil, RedisValue::Int(1)]),
"redis xautoclaim deleted_ids expected a redis array payload",
),
(
RedisValue::Array(vec![
bulk("0-0"),
RedisValue::Nil,
RedisValue::Array(vec![RedisValue::Nil]),
]),
"redis xautoclaim deleted_ids was not a string-compatible",
),
] {
let error = parse_reclaim_result(reply).expect_err("invalid reclaim reply");
let DataLayerError::UnexpectedValue(message) = error else {
panic!("reclaim parse error must keep its classification: {error}");
};
assert!(
message.starts_with(context),
"expected {context}, got {message}"
);
}
// Unlike read-group's filter, reclaim has always rejected an invalid value even if
// a later duplicate would overwrite it. Keep that validation order.
let reply = reclaim_reply(
bulk("1-0"),
fields_reply(
vec![
(bulk("payload"), RedisValue::Nil),
(bulk("payload"), bulk("later valid payload")),
],
false,
),
);
assert!(matches!(
parse_reclaim_result(reply),
Err(DataLayerError::UnexpectedValue(_))
));
}
@@ -0,0 +1,860 @@
use super::*;
use std::collections::BTreeSet;
use std::future::{poll_fn, Future};
use std::pin::Pin;
use std::task::Poll;
type TestConnection = ::redis::aio::MultiplexedConnection;
type ReadResult = Result<Vec<RuntimeQueueEntry>, DataLayerError>;
const TEST_GROUP: &str = "receive-workers";
const OWNER_BLOCK_MS: u64 = 60_000;
fn receive_lane_count() -> usize {
std::thread::available_parallelism()
.map(|value| value.get())
.unwrap_or(4)
.clamp(4, 16)
}
async fn receive_runtime(
protocol: &str,
command_timeout_ms: u64,
) -> Option<(TestRedisServer, RuntimeState, TestConnection)> {
let Some(server) = TestRedisServer::start().await else {
eprintln!(
"stream receive {protocol} skipped: isolated Redis fixture unavailable; check AETHER_REDIS_SERVER_BIN"
);
return None;
};
let mut admin = ::redis::Client::open(server.redis_url.clone())
.expect("test admin client")
.get_multiplexed_async_connection()
.await
.expect("test admin connection");
::redis::cmd("ACL")
.arg("SETUSER")
.arg("stream-reader")
.arg("on")
.arg(">stream-test-password")
.arg("~*")
.arg("+@all")
.query_async::<()>(&mut admin)
.await
.expect("test stream user");
let runtime = RuntimeState::redis_with_blocking_stream_lanes(
RedisClientConfig {
url: format!(
"redis://stream-reader:[email protected]:{}/7?protocol={protocol}",
server.port
),
key_prefix: Some(format!("receive-{protocol}")),
},
Some(command_timeout_ms),
Some(4),
)
.await
.expect("authenticated receive runtime in database 7");
::redis::cmd("SELECT")
.arg(7)
.query_async::<()>(&mut admin)
.await
.expect("admin selects test database");
eprintln!(
"stream receive fixture ready: protocol={protocol} db=7 port={} authenticated=true",
server.port
);
Some((server, runtime, admin))
}
async fn receive_group(runtime: &RuntimeState, stream: &str) {
RuntimeQueueStore::ensure_consumer_group(runtime, stream, TEST_GROUP, "0-0")
.await
.expect("receive consumer group");
}
fn receive_fields(sequence: usize) -> BTreeMap<String, String> {
BTreeMap::from([
(
"payload".to_string(),
format!("record-{sequence}\r\n\"quoted\"\\\u{4e2d}\u{6587}"),
),
("sequence".to_string(), sequence.to_string()),
("legacy_field".to_string(), "preserve exactly".to_string()),
])
}
async fn append_receive(runtime: &RuntimeState, stream: &str, sequence: usize) -> String {
RuntimeQueueStore::append_fields_with_maxlen(runtime, stream, &receive_fields(sequence), None)
.await
.expect("append receive entry")
}
async fn client_rows(admin: &mut TestConnection) -> Vec<BTreeMap<String, String>> {
let value = ::redis::cmd("CLIENT")
.arg("LIST")
.query_async::<String>(admin)
.await
.expect("Redis client list");
value
.lines()
.map(|line| {
line.split_whitespace()
.filter_map(|field| field.split_once('='))
.map(|(key, value)| (key.to_string(), value.to_string()))
.collect()
})
.collect()
}
fn blocked_rows(rows: &[BTreeMap<String, String>]) -> Vec<&BTreeMap<String, String>> {
rows.iter()
.filter(|row| row.get("flags").is_some_and(|flags| flags.contains('b')))
.collect()
}
async fn wait_for_blocked(admin: &mut TestConnection, expected: usize) {
tokio::time::timeout(Duration::from_secs(10), async {
loop {
if blocked_rows(&client_rows(admin).await).len() == expected {
return;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("Redis blocked clients reach expected count");
}
fn spawn_owner(
owners: &mut tokio::task::JoinSet<ReadResult>,
runtime: &RuntimeState,
stream: &'static str,
index: usize,
) {
let runtime = runtime.clone();
owners.spawn(async move {
RuntimeQueueStore::read_group(
&runtime,
stream,
TEST_GROUP,
&format!("owner-{index}"),
1,
Some(OWNER_BLOCK_MS),
)
.await
});
}
async fn abort_owners(owners: &mut tokio::task::JoinSet<ReadResult>) {
owners.abort_all();
while let Some(result) = owners.join_next().await {
assert!(result.expect_err("owner must be cancelled").is_cancelled());
}
}
async fn assert_receive_pending<F: Future + ?Sized>(mut future: Pin<&mut F>) {
poll_fn(|context| {
assert!(future.as_mut().poll(context).is_pending());
Poll::Ready(())
})
.await;
}
async fn pending_consumers(
admin: &mut TestConnection,
stream: &str,
) -> Vec<(String, String, u64, u64)> {
::redis::cmd("XPENDING")
.arg(stream)
.arg(TEST_GROUP)
.arg("-")
.arg("+")
.arg(100)
.query_async(admin)
.await
.expect("pending entry ownership")
}
#[tokio::test]
async fn redis_stream_receive_full_pool_waits_without_sending_and_cancellation_preserves_pel() {
for protocol in ["resp2", "resp3"] {
let Some((_server, runtime, mut admin)) = receive_runtime(protocol, 5_000).await else {
return;
};
let stream = "receive:blocked";
let fast_stream = "receive:nonblocking";
receive_group(&runtime, stream).await;
receive_group(&runtime, fast_stream).await;
let lanes = receive_lane_count();
let mut owners = tokio::task::JoinSet::new();
for index in 0..lanes {
spawn_owner(&mut owners, &runtime, stream, index);
}
wait_for_blocked(&mut admin, lanes).await;
let before = client_rows(&mut admin).await;
let owner_input_bytes = blocked_rows(&before)
.into_iter()
.map(|row| (row["id"].clone(), row.get("tot-net-in").cloned()))
.collect::<BTreeMap<_, _>>();
let mut waiter = Box::pin(RuntimeQueueStore::read_group(
&runtime,
stream,
TEST_GROUP,
"cancelled-waiter",
1,
Some(OWNER_BLOCK_MS),
));
assert_receive_pending(waiter.as_mut()).await;
// Complete unrelated round trips while the waiter remains polled and the owners block.
tokio::time::timeout(Duration::from_secs(5), async {
runtime.kv_set("receive-fast", "ready", None).await.unwrap();
assert_eq!(
runtime.kv_get("receive-fast").await.unwrap().as_deref(),
Some("ready")
);
let expected_id = append_receive(&runtime, fast_stream, 7).await;
let entries = RuntimeQueueStore::read_group(
&runtime,
fast_stream,
TEST_GROUP,
"nonblocking-reader",
1,
None,
)
.await
.unwrap();
assert_eq!(entries.len(), 1);
assert_eq!(entries[0].id, expected_id);
assert_eq!(entries[0].fields, receive_fields(7));
})
.await
.expect("full blocking pool must not delay fast or nonblocking stream lanes");
for _ in 0..3 {
tokio::task::yield_now().await;
let rows = client_rows(&mut admin).await;
let blocked = blocked_rows(&rows);
assert_eq!(blocked.len(), lanes);
for row in blocked {
assert_eq!(
row.get("tot-net-in"),
owner_input_bytes[&row["id"]].as_ref()
);
assert_eq!(row["qbuf"], "0", "waiter must not be sent behind a BLOCK");
}
assert_receive_pending(waiter.as_mut()).await;
}
drop(waiter);
assert_eq!(blocked_rows(&client_rows(&mut admin).await).len(), lanes);
assert!(pending_consumers(&mut admin, stream).await.is_empty());
abort_owners(&mut owners).await;
wait_for_blocked(&mut admin, 0).await;
let expected_id = append_receive(&runtime, stream, 8).await;
let entries = RuntimeQueueStore::read_group(
&runtime,
stream,
TEST_GROUP,
"replacement-reader",
1,
Some(100),
)
.await
.expect("replacement blocking connection");
assert_eq!(entries.len(), 1);
assert_eq!(entries[0].id, expected_id);
assert_eq!(entries[0].fields, receive_fields(8));
let pending = pending_consumers(&mut admin, stream).await;
assert_eq!(pending.len(), 1);
assert_eq!(pending[0].0, expected_id);
assert_eq!(pending[0].1, "replacement-reader");
::redis::cmd("SELECT")
.arg(0)
.query_async::<()>(&mut admin)
.await
.unwrap();
let other_database_len = ::redis::cmd("XLEN")
.arg(stream)
.query_async::<u64>(&mut admin)
.await
.unwrap();
assert_eq!(
other_database_len, 0,
"AUTH and SELECT must survive replacement connections"
);
}
}
#[tokio::test]
async fn redis_stream_receive_fast_consumer_reuses_free_lane_while_other_lanes_block() {
let Some((_server, runtime, mut admin)) = receive_runtime("resp2", 2_000).await else {
return;
};
let slow_stream = "receive:slow";
let fast_stream = "receive:ready";
receive_group(&runtime, slow_stream).await;
receive_group(&runtime, fast_stream).await;
let lanes = receive_lane_count();
let mut owners = tokio::task::JoinSet::new();
for index in 0..lanes - 1 {
spawn_owner(&mut owners, &runtime, slow_stream, index);
}
wait_for_blocked(&mut admin, lanes - 1).await;
let mut reader_connection_id = None;
for sequence in 0..lanes * 2 {
let expected_id = append_receive(&runtime, fast_stream, sequence).await;
let entries = RuntimeQueueStore::read_group(
&runtime,
fast_stream,
TEST_GROUP,
"fast-reader",
1,
Some(100),
)
.await
.expect("a free lane must remain reusable instead of rotating into a blocked lane");
assert_eq!(entries.len(), 1);
assert_eq!(entries[0].id, expected_id);
assert_eq!(entries[0].fields, receive_fields(sequence));
let rows = client_rows(&mut admin).await;
let idle_readers = rows
.iter()
.filter(|row| {
row.get("cmd").map(String::as_str) == Some("xreadgroup")
&& row.get("flags").is_some_and(|flags| !flags.contains('b'))
})
.collect::<Vec<_>>();
assert_eq!(idle_readers.len(), 1);
let current_id = &idle_readers[0]["id"];
if let Some(previous_id) = reader_connection_id.as_ref() {
assert_eq!(
current_id, previous_id,
"successful reads must reuse the same free connection"
);
} else {
reader_connection_id = Some(current_id.clone());
}
}
assert_eq!(
blocked_rows(&client_rows(&mut admin).await).len(),
lanes - 1
);
abort_owners(&mut owners).await;
wait_for_blocked(&mut admin, 0).await;
}
#[tokio::test]
async fn redis_stream_receive_timeout_discards_inflight_connection_before_reuse() {
let Some((_server, runtime, mut admin)) = receive_runtime("resp3", 1_000).await else {
return;
};
let stream = "receive:timeout";
receive_group(&runtime, stream).await;
let initial_connection_ids = client_rows(&mut admin)
.await
.into_iter()
.map(|row| row["id"].clone())
.collect::<BTreeSet<_>>();
// XREADGROUP is a write command; pausing writes makes its network response exceed the
// normal BLOCK-plus-grace timeout while read-only CLIENT diagnostics remain available.
::redis::cmd("CLIENT")
.arg("PAUSE")
.arg(30_000)
.arg("WRITE")
.query_async::<()>(&mut admin)
.await
.expect("pause test Redis writes");
let result = tokio::time::timeout(
Duration::from_secs(10),
RuntimeQueueStore::read_group(
&runtime,
stream,
TEST_GROUP,
"timed-out-reader",
1,
Some(100),
),
)
.await
.expect("read reaches its configured command deadline");
assert!(matches!(result, Err(DataLayerError::TimedOut(_))));
tokio::time::timeout(Duration::from_secs(10), async {
loop {
let connection_ids = client_rows(&mut admin)
.await
.into_iter()
.map(|row| row["id"].clone())
.collect::<BTreeSet<_>>();
if connection_ids.len() + 1 == initial_connection_ids.len()
&& connection_ids.is_subset(&initial_connection_ids)
{
return;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("timed-out connection must disconnect while its old command is still paused");
::redis::cmd("CLIENT")
.arg("UNPAUSE")
.query_async::<()>(&mut admin)
.await
.unwrap();
let expected_id = append_receive(&runtime, stream, 9).await;
let entries =
RuntimeQueueStore::read_group(&runtime, stream, TEST_GROUP, "after-timeout", 1, Some(100))
.await
.expect("replacement read after timeout");
assert_eq!(entries.len(), 1);
assert_eq!(entries[0].id, expected_id);
assert_eq!(entries[0].fields, receive_fields(9));
let pending = pending_consumers(&mut admin, stream).await;
assert_eq!(pending.len(), 1);
assert_eq!(pending[0].0, expected_id);
assert_eq!(pending[0].1, "after-timeout");
}
async fn set_pending_idle(admin: &mut TestConnection, stream: &str, ids: &[String]) {
let mut command = ::redis::cmd("XCLAIM");
command
.arg(stream)
.arg(TEST_GROUP)
.arg("initial-reader")
.arg(0);
for id in ids {
command.arg(id);
}
command.arg("IDLE").arg(120_000).arg("JUSTID");
let changed = command
.query_async::<Vec<String>>(admin)
.await
.expect("set pending idle");
assert_eq!(changed, ids);
}
#[tokio::test]
async fn redis_stream_receive_reclaim_pages_advance_past_fresh_prefix_and_deleted_entries() {
for protocol in ["resp2", "resp3"] {
let Some((_server, runtime, mut admin)) = receive_runtime(protocol, 5_000).await else {
return;
};
let stream = "receive:reclaim-pages";
receive_group(&runtime, stream).await;
let mut ids = Vec::new();
for sequence in 0..28 {
ids.push(append_receive(&runtime, stream, sequence).await);
}
let entries =
RuntimeQueueStore::read_group(&runtime, stream, TEST_GROUP, "initial-reader", 28, None)
.await
.expect("seed pending entries");
assert_eq!(entries.len(), ids.len());
set_pending_idle(&mut admin, stream, &ids[25..]).await;
assert_eq!(
RuntimeQueueStore::delete(&runtime, stream, &ids[26..27])
.await
.unwrap(),
1
);
let config = RuntimeQueueReclaimConfig {
min_idle_ms: 60_000,
count: 2,
};
let first = RuntimeQueueStore::claim_stale_page(
&runtime,
stream,
TEST_GROUP,
"reclaim-reader",
"0-0",
config,
)
.await
.expect("first reclaim page");
assert!(
first.entries.is_empty(),
"fresh prefix exceeds COUNT * 10 scan budget"
);
assert!(first.deleted_ids.is_empty());
assert_ne!(
first.next_start_id, "0-0",
"empty page must preserve continuation"
);
let mut cursor = first.next_start_id;
let mut reclaimed = BTreeMap::new();
let mut deleted = BTreeSet::new();
for _ in 0..8 {
let page = RuntimeQueueStore::claim_stale_page(
&runtime,
stream,
TEST_GROUP,
"reclaim-reader",
&cursor,
config,
)
.await
.expect("continued reclaim page");
assert!(page.entries.len() <= config.count);
for entry in page.entries {
assert!(reclaimed.insert(entry.id, entry.fields).is_none());
}
deleted.extend(page.deleted_ids);
cursor = page.next_start_id;
if cursor == "0-0" {
break;
}
}
assert_eq!(cursor, "0-0", "scan must eventually wrap");
assert_eq!(
reclaimed,
BTreeMap::from([
(ids[25].clone(), receive_fields(25)),
(ids[27].clone(), receive_fields(27))
])
);
assert_eq!(deleted, BTreeSet::from([ids[26].clone()]));
let pending = pending_consumers(&mut admin, stream).await;
assert_eq!(pending.len(), 27);
assert!(!pending.iter().any(|entry| entry.0 == ids[26]));
assert!(pending
.iter()
.filter(|entry| entry.1 == "reclaim-reader")
.all(|entry| entry.0 == ids[25] || entry.0 == ids[27]));
set_pending_idle(&mut admin, stream, &ids[..1]).await;
let restarted = RuntimeQueueStore::claim_stale_page(
&runtime,
stream,
TEST_GROUP,
"rescan-reader",
"0-0",
config,
)
.await
.expect("restart scan after cursor wraps");
assert_eq!(restarted.entries.len(), 1);
assert_eq!(restarted.entries[0].id, ids[0]);
assert_eq!(restarted.entries[0].fields, receive_fields(0));
assert!(restarted.deleted_ids.is_empty());
}
}
async fn interrupted_reclaim_preserves_pending(protocol: &str, cancel: bool) {
let command_timeout_ms = if cancel { 10_000 } else { 1_000 };
let Some((_server, runtime, mut admin)) = receive_runtime(protocol, command_timeout_ms).await
else {
return;
};
let stream = "receive:interrupted-reclaim";
receive_group(&runtime, stream).await;
let mut ids = Vec::new();
for sequence in 0..3 {
ids.push(append_receive(&runtime, stream, sequence).await);
}
assert_eq!(
RuntimeQueueStore::read_group(&runtime, stream, TEST_GROUP, "initial-reader", 3, None,)
.await
.unwrap()
.len(),
3
);
set_pending_idle(&mut admin, stream, &ids).await;
RuntimeQueueStore::delete(&runtime, stream, &ids[1..2])
.await
.unwrap();
let before = client_rows(&mut admin).await;
let initial_ids = before
.iter()
.map(|row| row["id"].clone())
.collect::<BTreeSet<_>>();
let input_bytes = before
.iter()
.filter(|row| row.get("user").map(String::as_str) == Some("stream-reader"))
.map(|row| {
(
row["id"].clone(),
row.get("tot-net-in")
.and_then(|value| value.parse::<u64>().ok()),
)
})
.collect::<BTreeMap<_, _>>();
::redis::cmd("CLIENT")
.arg("PAUSE")
.arg(30_000)
.arg("WRITE")
.query_async::<()>(&mut admin)
.await
.unwrap();
let config = RuntimeQueueReclaimConfig {
min_idle_ms: 60_000,
count: 1,
};
let claim_runtime = runtime.clone();
let claim = tokio::spawn(async move {
RuntimeQueueStore::claim_stale_page(
&claim_runtime,
stream,
TEST_GROUP,
"interrupted-reader",
"0-0",
config,
)
.await
});
if cancel {
tokio::time::timeout(Duration::from_secs(5), async {
loop {
let sent = client_rows(&mut admin).await.iter().any(|row| {
let Some(previous) = input_bytes.get(&row["id"]) else {
return false;
};
match (
previous,
row.get("tot-net-in")
.and_then(|value| value.parse::<u64>().ok()),
) {
(Some(previous), Some(current)) => current > *previous,
_ => row.get("cmd").map(String::as_str) == Some("xautoclaim"),
}
});
if sent {
return;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("claim reaches Redis before caller cancellation");
claim.abort();
assert!(claim.await.expect_err("cancelled reclaim").is_cancelled());
} else {
let result = tokio::time::timeout(Duration::from_secs(10), claim)
.await
.expect("reclaim command deadline")
.expect("reclaim task");
assert!(matches!(result, Err(DataLayerError::TimedOut(_))));
}
tokio::time::timeout(Duration::from_secs(10), async {
loop {
let remaining = client_rows(&mut admin)
.await
.into_iter()
.map(|row| row["id"].clone())
.collect::<BTreeSet<_>>();
if remaining.len() + 1 == initial_ids.len() && remaining.is_subset(&initial_ids) {
return;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("interrupted claim closes its socket before writes resume");
let pending = pending_consumers(&mut admin, stream).await;
assert_eq!(pending.len(), 3);
assert!(pending.iter().all(|entry| entry.1 == "initial-reader"));
::redis::cmd("CLIENT")
.arg("UNPAUSE")
.query_async::<()>(&mut admin)
.await
.unwrap();
let mut cursor = "0-0".to_string();
let mut recovered = BTreeMap::new();
let mut deleted = BTreeSet::new();
let mut successful_claims = 0;
for _ in 0..5 {
let page = RuntimeQueueStore::claim_stale_page(
&runtime,
stream,
TEST_GROUP,
"recovery-reader",
&cursor,
config,
)
.await
.expect("reclaim recovers after interrupted connection");
successful_claims += 1;
for entry in page.entries {
assert!(recovered.insert(entry.id, entry.fields).is_none());
}
deleted.extend(page.deleted_ids);
cursor = page.next_start_id;
if cursor == "0-0" {
break;
}
}
assert_eq!(cursor, "0-0");
assert_eq!(
recovered,
BTreeMap::from([
(ids[0].clone(), receive_fields(0)),
(ids[2].clone(), receive_fields(2))
])
);
assert_eq!(deleted, BTreeSet::from([ids[1].clone()]));
let pending = pending_consumers(&mut admin, stream).await;
assert_eq!(pending.len(), 2);
assert!(pending.iter().all(|entry| entry.1 == "recovery-reader"));
let diagnostics = runtime.redis_diagnostics().await.unwrap().unwrap();
let lane = diagnostics
.lanes
.iter()
.find(|lane| lane.lane == "blocking_stream")
.unwrap();
assert_eq!(lane.command_timeouts, u64::from(!cancel));
assert!(lane.command_count >= successful_claims);
}
#[tokio::test]
async fn redis_stream_receive_reclaim_cancellation_closes_connection_and_preserves_pending() {
interrupted_reclaim_preserves_pending("resp2", true).await;
}
#[tokio::test]
async fn redis_stream_receive_reclaim_timeout_closes_connection_and_preserves_pending() {
interrupted_reclaim_preserves_pending("resp3", false).await;
}
#[tokio::test]
async fn redis_stream_receive_reclaim_waits_for_read_lease_and_continues_after_release() {
let Some((_server, runtime, mut admin)) = receive_runtime("resp2", 1_000).await else {
return;
};
let blocked_stream = "receive:reclaim-pool-blocked";
let pending_stream = "receive:reclaim-pool-pending";
receive_group(&runtime, blocked_stream).await;
receive_group(&runtime, pending_stream).await;
let ids = vec![
append_receive(&runtime, pending_stream, 0).await,
append_receive(&runtime, pending_stream, 1).await,
];
assert_eq!(
RuntimeQueueStore::read_group(
&runtime,
pending_stream,
TEST_GROUP,
"initial-reader",
2,
None
)
.await
.unwrap()
.len(),
2
);
set_pending_idle(&mut admin, pending_stream, &ids).await;
let lanes = receive_lane_count();
let mut owners = tokio::task::JoinSet::new();
for index in 0..lanes - 1 {
spawn_owner(&mut owners, &runtime, blocked_stream, index);
}
let owner_runtime = runtime.clone();
let release_owner = owners.spawn(async move {
RuntimeQueueStore::read_group(
&owner_runtime,
blocked_stream,
TEST_GROUP,
"released-owner",
1,
Some(OWNER_BLOCK_MS),
)
.await
});
wait_for_blocked(&mut admin, lanes).await;
let config = RuntimeQueueReclaimConfig {
min_idle_ms: 60_000,
count: 1,
};
let result = tokio::time::timeout(
Duration::from_secs(10),
RuntimeQueueStore::claim_stale_page(
&runtime,
pending_stream,
TEST_GROUP,
"timed-out-waiter",
"0-0",
config,
),
)
.await
.expect("checkout uses the original command deadline");
assert!(matches!(result, Err(DataLayerError::TimedOut(_))));
assert_eq!(blocked_rows(&client_rows(&mut admin).await).len(), lanes);
assert!(pending_consumers(&mut admin, pending_stream)
.await
.iter()
.all(|entry| entry.1 == "initial-reader"));
let mut cancelled_claim = Box::pin(RuntimeQueueStore::claim_stale_page(
&runtime,
pending_stream,
TEST_GROUP,
"cancelled-waiter",
"0-0",
config,
));
assert_receive_pending(cancelled_claim.as_mut()).await;
drop(cancelled_claim);
assert_eq!(blocked_rows(&client_rows(&mut admin).await).len(), lanes);
let mut claim = Box::pin(RuntimeQueueStore::claim_stale_page(
&runtime,
pending_stream,
TEST_GROUP,
"after-release",
"0-0",
config,
));
assert_receive_pending(claim.as_mut()).await;
release_owner.abort();
assert!(owners
.join_next()
.await
.unwrap()
.expect_err("released owner cancelled")
.is_cancelled());
let first = tokio::time::timeout(Duration::from_secs(5), claim)
.await
.expect("claim receives released pool capacity")
.expect("claim succeeds after read releases lease");
assert_eq!(first.entries.len(), 1);
assert_eq!(first.entries[0].id, ids[0]);
assert_eq!(first.entries[0].fields, receive_fields(0));
assert_ne!(first.next_start_id, "0-0");
let next_id = append_receive(&runtime, pending_stream, 2).await;
let next_read = RuntimeQueueStore::read_group(
&runtime,
pending_stream,
TEST_GROUP,
"read-after-claim",
1,
Some(100),
)
.await
.expect("completed claim returns its lease for the next read");
assert_eq!(next_read.len(), 1);
assert_eq!(next_read[0].id, next_id);
let second = RuntimeQueueStore::claim_stale_page(
&runtime,
pending_stream,
TEST_GROUP,
"after-release",
&first.next_start_id,
config,
)
.await
.expect("completed read returns its lease for the next claim");
assert_eq!(second.entries.len(), 1);
assert_eq!(second.entries[0].id, ids[1]);
assert_eq!(second.entries[0].fields, receive_fields(1));
assert_eq!(
blocked_rows(&client_rows(&mut admin).await).len(),
lanes - 1
);
abort_owners(&mut owners).await;
wait_for_blocked(&mut admin, 0).await;
}
@@ -0,0 +1,175 @@
use std::sync::OnceLock;
use super::client::RedisBlockingStreamLease;
use super::runtime::USAGE_LIMIT_CHECK_AND_CONSUME_SCRIPT;
use super::{cmd, RedisConnectionLane, RedisConnectionRouter};
use crate::error::RedisResultExt;
use crate::{DataLayerError, UsageLimitInput};
const COPY_CHUNK_SIZE: i64 = 512;
const MAX_COPY_ATTEMPTS: usize = 8;
const COPY_SCRIPT: &str = include_str!("usage_copy.lua");
const COMMIT_PREFIX: &str = include_str!("usage_copy_commit.lua");
fn usage_args(command: &mut redis::Cmd, input: &UsageLimitInput<'_>) {
command.arg(input.now_unix_ms).arg(input.event_id);
for rule in input.rules {
command
.arg(rule.limit)
.arg(rule.window_seconds)
.arg(rule.retention_seconds);
}
}
pub(super) async fn check_and_consume(
connections: &RedisConnectionRouter,
keys: &[String],
input: &UsageLimitInput<'_>,
) -> Result<Vec<i64>, DataLayerError> {
static SCRIPT: OnceLock<redis::Script> = OnceLock::new();
let script = SCRIPT.get_or_init(|| redis::Script::new(USAGE_LIMIT_CHECK_AND_CONSUME_SCRIPT));
let mut invocation = script.prepare_invoke();
for key in keys {
invocation.key(key);
}
invocation.arg(input.now_unix_ms).arg(input.event_id);
for rule in input.rules {
invocation
.arg(rule.limit)
.arg(rule.window_seconds)
.arg(rule.retention_seconds);
}
let result: Vec<i64> = invocation
.invoke_async(&mut connections.connection(RedisConnectionLane::Fast))
.await
.map_redis_err()?;
if result.first() != Some(&2) {
return Ok(result);
}
let mut lease = connections.usage_cleanup_connection().await?;
for _ in 0..MAX_COPY_ATTEMPTS {
let mut temporary_keys = Vec::new();
let result = copy_and_commit(&mut lease, keys, input, &mut temporary_keys).await;
// EXEC clears WATCH even on a conflict. Errors discard the lease,
// including any pending WATCH/MULTI state.
if let Ok(Some((result, committed))) = result.as_ref() {
// Never add another fallible round trip after admission. EXEC already
// cleared WATCH; the recheck fast path instead discards its watched lease.
if *committed {
lease.recycle();
}
return Ok(result.clone());
} else if result.is_ok() {
lease.query(&cmd("UNWATCH")).await?;
if !temporary_keys.is_empty() {
lease.query(cmd("UNLINK").arg(&temporary_keys)).await?;
}
} else if !temporary_keys.is_empty() {
// A fresh connection cannot accidentally queue cleanup inside a failed MULTI.
let mut connection = connections.connection(RedisConnectionLane::Admin);
let _ = cmd("UNLINK")
.arg(&temporary_keys)
.query_async::<usize>(&mut connection)
.await;
}
result?;
tokio::task::yield_now().await;
}
lease.recycle();
Err(DataLayerError::Redis(
"usage window cleanup conflicted repeatedly; retry the request".to_string(),
))
}
async fn copy_and_commit(
lease: &mut RedisBlockingStreamLease,
keys: &[String],
input: &UsageLimitInput<'_>,
temporary_keys: &mut Vec<String>,
) -> Result<Option<(Vec<i64>, bool)>, DataLayerError> {
lease.query(cmd("WATCH").arg(keys)).await?;
let mut check = cmd("EVAL");
check
.arg(USAGE_LIMIT_CHECK_AND_CONSUME_SCRIPT)
.arg(keys.len())
.arg(keys);
usage_args(&mut check, input);
let plan: Vec<i64> =
redis::from_owned_redis_value(lease.query(&check).await?).map_redis_err()?;
if plan.first() != Some(&2) {
return Ok(Some((plan, false)));
}
if plan.len() < 4 || (plan.len() - 1) % 3 != 0 {
return Err(DataLayerError::UnexpectedValue(
"invalid usage window copy plan".to_string(),
));
}
let nonce = uuid::Uuid::new_v4();
for window in plan[1..].chunks_exact(3) {
let [index, expired, live] = [window[0], window[1], window[2]];
let key = usize::try_from(index - 1)
.ok()
.and_then(|index| keys.get(index))
.filter(|_| expired > 0 && live > 0)
.ok_or_else(|| {
DataLayerError::UnexpectedValue("invalid usage window copy range".to_string())
})?;
let temporary = format!("{key}:__usage_copy:{nonce}");
let mut offset = 0;
while offset < live {
let take = COPY_CHUNK_SIZE.min(live - offset);
let value = lease
.query(
cmd("EVAL")
.arg(COPY_SCRIPT)
.arg(2)
.arg(key)
.arg(&temporary)
.arg(expired + offset)
.arg(expired + offset + take - 1)
.arg(offset),
)
.await?;
let copied: i64 = redis::from_owned_redis_value(value).map_redis_err()?;
if offset == 0 && copied > 0 {
temporary_keys.push(temporary.clone());
}
if copied != take {
return Ok(None);
}
offset += take;
}
}
static COMMIT: OnceLock<String> = OnceLock::new();
let source =
COMMIT.get_or_init(|| format!("{COMMIT_PREFIX}\n{USAGE_LIMIT_CHECK_AND_CONSUME_SCRIPT}"));
let mut commit = cmd("EVAL");
commit
.arg(source)
.arg(keys.len() + temporary_keys.len())
.arg(keys)
.arg(&*temporary_keys)
.arg(keys.len());
usage_args(&mut commit, input);
for window in plan[1..].chunks_exact(3) {
commit.arg(window[0]).arg(window[2]);
}
lease.query(&cmd("MULTI")).await?;
lease.query(&commit).await?;
let replies: Option<Vec<redis::Value>> =
redis::from_owned_redis_value(lease.query(&cmd("EXEC")).await?).map_redis_err()?;
let Some(mut replies) = replies else {
return Ok(None);
};
let reply = replies.pop().ok_or_else(|| {
DataLayerError::UnexpectedValue("empty usage window commit response".to_string())
})?;
let result: Vec<i64> = redis::from_owned_redis_value(reply).map_redis_err()?;
if result.first() == Some(&2) {
return Ok(None);
}
Ok(Some((result, true)))
}
@@ -0,0 +1,22 @@
local source, target = KEYS[1], KEYS[2]
local offset = tonumber(ARGV[3])
if offset == 0 and redis.call('EXISTS', target) ~= 0 then
return redis.error_reply('usage copy temporary key already exists')
end
if offset > 0 and redis.call('ZCARD', target) ~= offset then return -1 end
local rows = redis.call('ZRANGE', source, ARGV[1], ARGV[2], 'WITHSCORES')
if #rows == 0 then return 0 end
if not redis.acl_check_cmd('PEXPIRE', target, 60000)
or not redis.acl_check_cmd('UNLINK', target) then
return redis.error_reply('usage copy temporary key permission denied')
end
local args = {}
for i = 1, #rows, 2 do
args[#args + 1] = rows[i + 1]
args[#args + 1] = rows[i]
end
-- TTL is established in the same command as the first allocation. A cancelled
-- caller or a process crash cannot leave a permanent scratch key behind.
redis.call('ZADD', target, unpack(args))
redis.call('PEXPIRE', target, 60000)
return #rows / 2
@@ -0,0 +1,35 @@
local rule_count = tonumber(table.remove(ARGV, 1))
local swaps = #KEYS - rule_count
local ttls = {}
-- WATCH covers every source, including modifications made by older instances.
-- Check every temporary key and every permission before replacing any source.
for i = 1, swaps do
local position = rule_count * 3 + 2 + (i - 1) * 2
local index = tonumber(ARGV[position + 1])
local live = tonumber(ARGV[position + 2])
local source, target = KEYS[index], KEYS[rule_count + i]
local ttl = redis.call('PTTL', source)
if ttl == 0 or ttl < -1 or redis.call('ZCARD', target) ~= live then return {2} end
if not redis.acl_check_cmd('UNLINK', source)
or not redis.acl_check_cmd('RENAME', target, source)
or not redis.acl_check_cmd('PEXPIRE', target, math.max(1, ttl))
or not redis.acl_check_cmd('PERSIST', target) then
return redis.error_reply('usage copy commit permission denied')
end
ttls[i] = ttl
end
for i = 1, swaps do
local position = rule_count * 3 + 2 + (i - 1) * 2
local index = tonumber(ARGV[position + 1])
local source, target = KEYS[index], KEYS[rule_count + i]
if ttls[i] > 0 then
redis.call('PEXPIRE', target, ttls[i])
else
redis.call('PERSIST', target)
end
redis.call('UNLINK', source)
redis.call('RENAME', target, source)
end
for i = #KEYS, rule_count + 1, -1 do KEYS[i] = nil end
ARGV[rule_count * 3 + 3] = 'inline'
-- The original prune/check/consume script follows in this same EXEC/EVAL.
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,25 @@
/// Maximum number of members examined by one server-side window aggregation.
pub const SCORE_WINDOW_AGGREGATION_MEMBER_LIMIT: usize = 512;
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub struct ScoreWindowU64Stats {
pub sum: u64,
pub positive_count: u64,
}
impl ScoreWindowU64Stats {
/// Members encode the value after their final colon. Invalid or zero values
/// do not contribute to the positive sample count.
pub fn from_members<'a>(members: impl IntoIterator<Item = &'a str>) -> Self {
let mut stats = Self::default();
for member in members {
let value = member
.rsplit_once(':')
.and_then(|(_, suffix)| suffix.parse::<u64>().ok())
.unwrap_or(0);
stats.sum = stats.sum.saturating_add(value);
stats.positive_count += u64::from(value > 0);
}
stats
}
}
@@ -720,6 +720,71 @@ mod tests {
.expect("candidate should build")
}
#[test]
fn runtime_projection_preserves_concurrency_rpm_and_failure_cooldown() {
let statuses = [
RequestCandidateStatus::Available,
RequestCandidateStatus::Unused,
RequestCandidateStatus::Pending,
RequestCandidateStatus::Streaming,
RequestCandidateStatus::Success,
RequestCandidateStatus::Failed,
RequestCandidateStatus::Cancelled,
RequestCandidateStatus::Skipped,
];
let mut candidates = Vec::new();
for (index, status) in statuses.into_iter().cycle().take(40).enumerate() {
let mut row = stored_candidate(&index.to_string(), status, 100 - index as i64);
row.api_key_id = Some("api-key".into());
row.concurrent_requests = Some(index as u32);
row.started_at_unix_ms = (index % 2 == 0).then_some(99_000);
row.finished_at_unix_ms = (index % 3 == 0).then_some(101_000);
row.extra_data = Some(serde_json::json!({"stream_completed": true}));
candidates.push(row);
}
let projected = candidates
.iter()
.map(StoredRequestCandidate::runtime_snapshot)
.collect::<Vec<_>>();
for now in [90, 101, 160, 401] {
for rows in [&candidates[..], &candidates[5..], &candidates[39..]] {
let slim = &projected[candidates.len() - rows.len()..];
assert_eq!(
count_recent_active_requests_for_api_key(rows, "api-key", now),
count_recent_active_requests_for_api_key(slim, "api-key", now)
);
assert_eq!(
count_recent_active_requests_for_provider(rows, "provider-a", now),
count_recent_active_requests_for_provider(slim, "provider-a", now)
);
assert_eq!(
count_recent_active_requests_for_provider_key(rows, "key-a", now),
count_recent_active_requests_for_provider_key(slim, "key-a", now)
);
assert_eq!(
count_recent_rpm_requests_for_provider_key_since(rows, "key-a", now, Some(80)),
count_recent_rpm_requests_for_provider_key_since(slim, "key-a", now, Some(80))
);
assert_eq!(
is_candidate_in_recent_failure_cooldown(
rows,
"provider-a",
"endpoint-a",
"key-a",
now
),
is_candidate_in_recent_failure_cooldown(
slim,
"provider-a",
"endpoint-a",
"key-a",
now
)
);
}
}
}
fn provider_catalog_key(id: &str) -> StoredProviderCatalogKey {
StoredProviderCatalogKey::new(
id.to_string(),
@@ -13,6 +13,7 @@ aether-crypto.workspace = true
aether-data.workspace = true
aether-data-contracts.workspace = true
aether-gateway = { workspace = true, features = ["testkit"] }
aether-runtime.workspace = true
aether-runtime-state.workspace = true
aether-testkit = { workspace = true, features = ["gateway", "postgres"] }
axum.workspace = true
@@ -18,7 +18,7 @@ use axum::http::StatusCode;
use axum::response::{IntoResponse, Response};
use axum::routing::any;
use axum::{extract::Request, Json, Router};
use futures_util::{SinkExt, StreamExt};
use futures_util::{stream::FuturesUnordered, SinkExt, StreamExt};
use reqwest::Method;
use serde::Serialize;
use serde_json::json;
@@ -83,6 +83,8 @@ struct CapacityCurvePointResult {
successful_requests: usize,
rejected_requests: usize,
failed_requests: usize,
status_counts: BTreeMap<u16, usize>,
non_success_status_samples: serde_json::Value,
throughput_rps: u64,
p50_ms: u64,
p95_ms: u64,
@@ -112,8 +114,16 @@ struct GateMetricSnapshot {
rejected_total: u64,
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
let runtime = tokio::runtime::Builder::new_multi_thread()
.enable_all()
.thread_stack_size(8 * 1024 * 1024)
.build()?;
runtime.block_on(run())
}
async fn run() -> Result<(), Box<dyn std::error::Error>> {
init_test_runtime_for("capacity-curve-baseline");
let config = parse_args(std::env::args().skip(1).collect())?;
let report = run_suite(&config).await?;
@@ -215,11 +225,11 @@ async fn run_gateway_curve(
.await
.map_err(std::io::Error::other)?;
let duration_ms = started_at.elapsed().as_millis() as u64;
let metrics = capture_gate_metrics(
&format!("{}/_gateway/metrics", gateway.base_url()),
gate_name,
)
.await?;
let samples = gateway
.metric_samples()
.await
.map_err(std::io::Error::other)?;
let metrics = gate_metrics(&samples, gate_name)?;
points.push(capacity_point(
*limit,
total_requests,
@@ -319,6 +329,21 @@ async fn run_tunnel_curve(
let peer = connect_protocol_peer(tunnel.base_url(), config.tunnel_hold).await?;
let total_requests =
total_requests_for_limit(relay_concurrency, config.requests_per_point_multiplier);
let envelope = relay_envelope();
let body_offset =
4 + u32::from_be_bytes(envelope[..4].try_into().expect("metadata length")) as usize;
verify_tunnel_fixture(&tunnel, &envelope, body_offset, config.timeout).await?;
let header_sets = (0..total_requests)
.map(|_| {
let mut headers =
tunnel.relay_headers(&envelope[..body_offset], &envelope[body_offset..]);
headers.insert(
"content-type".to_string(),
"application/octet-stream".to_string(),
);
headers
})
.collect();
let probe = HttpLoadProbeConfig {
url: format!(
"{tunnel_base}{TUNNEL_RELAY_PATH_PREFIX}/node-baseline",
@@ -329,7 +354,8 @@ async fn run_tunnel_curve(
"content-type".to_string(),
"application/octet-stream".to_string(),
)]),
body: Some(relay_envelope()),
header_sets,
body: Some(envelope),
total_requests,
concurrency: relay_concurrency,
timeout: config.timeout,
@@ -350,7 +376,8 @@ async fn run_tunnel_curve(
result,
metrics,
));
drop(peer);
peer.abort();
let _ = peer.await;
}
Ok(CapacityCurveScenarioReport {
@@ -362,6 +389,58 @@ async fn run_tunnel_curve(
})
}
async fn verify_tunnel_fixture(
tunnel: &TunnelHarness,
envelope: &[u8],
body_offset: usize,
timeout: Duration,
) -> Result<(), Box<dyn std::error::Error>> {
let client = reqwest::Client::builder().timeout(timeout).build()?;
let url = format!(
"{}{TUNNEL_RELAY_PATH_PREFIX}/{TUNNEL_HARNESS_NODE_ID}",
tunnel.base_url()
);
let unsigned = client.post(&url).body(envelope.to_vec()).send().await?;
if unsigned.status() != StatusCode::FORBIDDEN {
return Err(std::io::Error::other("unsigned tunnel preflight was not rejected").into());
}
let mut signed = client.post(&url).body(envelope.to_vec());
for (name, value) in tunnel.relay_headers(&envelope[..body_offset], &envelope[body_offset..]) {
signed = signed.header(name, value);
}
let signed = signed.build()?;
let mut tampered = signed
.try_clone()
.expect("buffered relay request should clone");
let mut tampered_body = envelope.to_vec();
*tampered_body
.last_mut()
.expect("relay body should be nonempty") ^= 1;
*tampered.body_mut() = Some(tampered_body.into());
if client.execute(tampered).await?.status() != StatusCode::FORBIDDEN {
return Err(std::io::Error::other("tampered tunnel preflight was not rejected").into());
}
let response = client
.execute(
signed
.try_clone()
.expect("buffered relay request should clone"),
)
.await?;
let status = response.status();
let body = response.text().await?;
if status != StatusCode::OK || body != "capacity-tunnel-stream" {
return Err(std::io::Error::other(format!(
"signed tunnel preflight failed: {status}: {body}"
))
.into());
}
if client.execute(signed).await?.status() != StatusCode::FORBIDDEN {
return Err(std::io::Error::other("replayed tunnel preflight was not rejected").into());
}
Ok(())
}
fn capacity_point(
limit: usize,
total_requests: usize,
@@ -395,7 +474,10 @@ fn capacity_point(
duration_ms,
successful_requests,
rejected_requests,
failed_requests: result.failed_requests,
failed_requests: total_requests.saturating_sub(successful_requests + rejected_requests),
status_counts: result.status_counts,
non_success_status_samples: serde_json::to_value(result.non_success_status_samples)
.expect("HTTP status samples should serialize"),
throughput_rps,
p50_ms: result.p50_ms,
p95_ms: result.p95_ms,
@@ -440,27 +522,29 @@ async fn capture_gate_metrics(
let samples = fetch_prometheus_samples(metrics_url)
.await
.map_err(std::io::Error::other)?;
gate_metrics(&samples, gate_name)
}
fn gate_metrics(
samples: &[aether_testkit::PrometheusSample],
gate_name: &str,
) -> Result<GateMetricSnapshot, Box<dyn std::error::Error>> {
let required = |name| {
find_metric_value_u64(samples, name, &[("gate", gate_name)])
.or_else(|| {
find_metric_value_u64(
samples,
&format!("aether_testkit_{name}"),
&[("gate", gate_name)],
)
})
.ok_or_else(|| std::io::Error::other(format!("missing {name} for gate {gate_name}")))
};
Ok(GateMetricSnapshot {
in_flight: find_metric_value_u64(&samples, "concurrency_in_flight", &[("gate", gate_name)])
.unwrap_or_default(),
available_permits: find_metric_value_u64(
&samples,
"concurrency_available_permits",
&[("gate", gate_name)],
)
.unwrap_or_default(),
high_watermark: find_metric_value_u64(
&samples,
"concurrency_high_watermark",
&[("gate", gate_name)],
)
.unwrap_or_default(),
rejected_total: find_metric_value_u64(
&samples,
"concurrency_rejected_total",
&[("gate", gate_name)],
)
.unwrap_or_default(),
in_flight: required("concurrency_in_flight")?,
available_permits: required("concurrency_available_permits")?,
high_watermark: required("concurrency_high_watermark")?,
rejected_total: required("concurrency_rejected_total")?,
})
}
@@ -700,25 +784,35 @@ async fn connect_protocol_peer(
))
.await?;
Ok(tokio::spawn(async move {
while let Some(message) = stream.next().await {
let Ok(message) = message else {
break;
};
match message {
Message::Binary(data)
if handle_binary_frame(&mut sink, data.to_vec(), hold)
.await
.is_err() =>
{
break;
let mut responses = FuturesUnordered::new();
loop {
tokio::select! {
message = stream.next() => {
match message {
Some(Ok(Message::Binary(data))) => {
match handle_binary_frame(&mut sink, data.to_vec()).await {
Ok(Some(stream_id)) => responses.push(async move {
tokio::time::sleep(hold).await;
stream_id
}),
Ok(None) => {},
Err(_) => break,
}
}
Some(Ok(Message::Ping(payload))) => {
if sink.send(Message::Pong(payload)).await.is_err() {
break;
}
}
None | Some(Err(_)) | Some(Ok(Message::Close(_))) => break,
_ => {},
}
}
Message::Ping(payload)
if sink.send(Message::Pong(payload.clone())).await.is_err() =>
{
break;
Some(stream_id) = responses.next(), if !responses.is_empty() => {
if send_protocol_response(&mut sink, stream_id).await.is_err() {
break;
}
}
Message::Close(_) => break,
_ => {}
}
}
let _ = sink.close().await;
@@ -728,13 +822,12 @@ async fn connect_protocol_peer(
async fn handle_binary_frame<S>(
sink: &mut S,
data: Vec<u8>,
hold: Duration,
) -> Result<(), tokio_tungstenite::tungstenite::Error>
) -> Result<Option<u32>, tokio_tungstenite::tungstenite::Error>
where
S: SinkExt<Message, Error = tokio_tungstenite::tungstenite::Error> + Unpin,
{
let Some(header) = protocol::FrameHeader::parse(&data) else {
return Ok(());
return Ok(None);
};
match header.msg_type {
protocol::PING => {
@@ -755,48 +848,57 @@ where
.await?;
}
if header.flags & protocol::FLAG_END_STREAM == 0 {
return Ok(());
return Ok(None);
}
tokio::time::sleep(hold).await;
let response_meta = protocol::ResponseMeta {
status: 200,
headers: vec![(
"content-type".to_string(),
"text/plain; charset=utf-8".to_string(),
)],
};
let response_meta_json =
serde_json::to_vec(&response_meta).expect("response metadata should serialize");
sink.send(Message::Binary(
protocol::encode_frame(
header.stream_id,
protocol::RESPONSE_HEADERS,
0,
&response_meta_json,
)
.into(),
))
.await?;
for chunk in [
b"capacity-".as_slice(),
b"tunnel-".as_slice(),
b"stream".as_slice(),
] {
sink.send(Message::Binary(
protocol::encode_frame(header.stream_id, protocol::RESPONSE_BODY, 0, chunk)
.into(),
))
.await?;
}
sink.send(Message::Binary(
protocol::encode_frame(header.stream_id, protocol::STREAM_END, 0, &[]).into(),
))
.await?;
return Ok(Some(header.stream_id));
}
_ => {}
}
Ok(None)
}
async fn send_protocol_response<S>(
sink: &mut S,
stream_id: u32,
) -> Result<(), tokio_tungstenite::tungstenite::Error>
where
S: SinkExt<Message, Error = tokio_tungstenite::tungstenite::Error> + Unpin,
{
let response_meta = protocol::ResponseMeta {
status: 200,
headers: vec![(
"content-type".to_string(),
"text/plain; charset=utf-8".to_string(),
)],
};
let response_meta_json =
serde_json::to_vec(&response_meta).expect("response metadata should serialize");
sink.send(Message::Binary(
protocol::encode_frame(
stream_id,
protocol::RESPONSE_HEADERS,
0,
&response_meta_json,
)
.into(),
))
.await?;
for chunk in [
b"capacity-".as_slice(),
b"tunnel-".as_slice(),
b"stream".as_slice(),
] {
sink.send(Message::Binary(
protocol::encode_frame(stream_id, protocol::RESPONSE_BODY, 0, chunk).into(),
))
.await?;
}
sink.send(Message::Binary(
protocol::encode_frame(stream_id, protocol::STREAM_END, 0, &[]).into(),
))
.await?;
Ok(())
}
@@ -141,8 +141,13 @@ impl SummaryCollector {
}
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
run()
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run() -> Result<(), Box<dyn std::error::Error>> {
init_test_runtime_for("dependency-pressure-baseline");
let config = parse_args(std::env::args().skip(1).collect())?;
let report = run_suite(&config).await?;
@@ -172,8 +172,13 @@ impl RecoveryCollector {
}
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
run()
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run() -> Result<(), Box<dyn std::error::Error>> {
init_test_runtime_for("failure-recovery-baseline");
let config = parse_args(std::env::args().skip(1).collect())?;
let report = run_suite(&config).await?;
@@ -7,6 +7,7 @@ use std::io;
use std::io::Write;
use std::path::{Path, PathBuf};
use aether_crypto::PythonFernetCompat;
use aether_data::repository::auth::CreateStandaloneApiKeyRecord;
use aether_data::repository::wallet::WalletLookupKey;
use aether_data::{
@@ -16,6 +17,7 @@ use aether_data_contracts::repository::global_models::{
CreateAdminGlobalModelRecord, UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
};
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyOAuthCredentialFence,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use serde_json::json;
@@ -200,6 +202,12 @@ impl Config {
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
let config = Config::from_env_and_args().map_err(|err| format!("invalid config: {err}"))?;
let encryption_key = env_value("AETHER_GATEWAY_DATA_ENCRYPTION_KEY")
.or_else(|| env_value("ENCRYPTION_KEY"))
.ok_or(
"set AETHER_GATEWAY_DATA_ENCRYPTION_KEY or ENCRYPTION_KEY to the gateway's encryption key before seeding",
)?;
let secret_cipher = PythonFernetCompat::from_secret(&encryption_key);
let backends = DataBackends::from_config(DataLayerConfig::from_database(SqlDatabaseConfig {
driver: DatabaseDriver::Postgres,
@@ -215,10 +223,10 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
},
}))?;
seed_provider_catalog(&backends, &config).await?;
seed_provider_catalog(&backends, &config, &secret_cipher).await?;
seed_models(&backends, &config).await?;
let operator_user_id = seed_operator_user(&backends, &config).await?;
seed_api_keys(&backends, &config, &operator_user_id).await?;
seed_api_keys(&backends, &config, &operator_user_id, &secret_cipher).await?;
verify_candidate_selection(&backends, &config).await?;
write_outputs(&config)?;
@@ -243,6 +251,7 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn seed_provider_catalog(
backends: &DataBackends,
config: &Config,
secret_cipher: &PythonFernetCompat,
) -> Result<(), Box<dyn std::error::Error>> {
let reader = backends
.read()
@@ -326,7 +335,7 @@ async fn seed_provider_catalog(
)?
.with_transport_fields(
Some(json!(["openai:chat"])),
Some(config.provider_api_key.clone()),
Some(secret_cipher.encrypt_plaintext(&config.provider_api_key)?),
None,
None,
None,
@@ -341,17 +350,40 @@ async fn seed_provider_catalog(
Some(json!({"openai:chat": {"state": "closed"}})),
);
if reader
.list_keys_by_ids(std::slice::from_ref(&config.provider_key_id))
.await?
.is_empty()
{
writer.create_key(&provider_key).await?;
} else {
writer.update_key(&provider_key).await?;
for _ in 0..8 {
let existing = reader
.list_keys_by_ids(std::slice::from_ref(&config.provider_key_id))
.await?
.into_iter()
.next();
let Some(existing) = existing else {
writer.create_key(&provider_key).await?;
return Ok(());
};
if existing.provider_id != config.provider_id {
return Err("existing pressure provider key belongs to a different provider".into());
}
// Randomized ciphertext changes on every seed. Fence against the observed
// credential and preserve runtime fields when rotating the configured key.
let update = ProviderCatalogKeyAdminCasUpdate {
expected_encrypted_auth_config: existing.encrypted_auth_config.clone(),
expected_credential: ProviderCatalogKeyOAuthCredentialFence {
encrypted_api_key: existing.encrypted_api_key,
auth_type: existing.auth_type,
provider_id: existing.provider_id,
provider_type: provider.provider_type.clone(),
},
key: provider_key.clone(),
codex_rotation: None,
reset_oauth_runtime: true,
};
if writer.compare_and_update_key_admin_state(&update).await? {
return Ok(());
}
}
Ok(())
Err("pressure provider key changed repeatedly during seed; retry initialization".into())
}
fn pressure_provider_transport_config(mock_upstream_h2c: bool) -> Option<serde_json::Value> {
@@ -466,9 +498,10 @@ async fn seed_api_keys(
backends: &DataBackends,
config: &Config,
operator_user_id: &str,
secret_cipher: &PythonFernetCompat,
) -> Result<(), Box<dyn std::error::Error>> {
for index in 0..config.api_key_count {
seed_api_key(backends, config, operator_user_id, index).await?;
seed_api_key(backends, config, operator_user_id, index, secret_cipher).await?;
}
Ok(())
}
@@ -478,6 +511,7 @@ async fn seed_api_key(
config: &Config,
operator_user_id: &str,
key_index: usize,
secret_cipher: &PythonFernetCompat,
) -> Result<(), Box<dyn std::error::Error>> {
let auth_reader = backends
.read()
@@ -494,17 +528,28 @@ async fn seed_api_key(
let api_key_id = pressure_api_key_id(config, key_index);
let api_key_value = pressure_api_key_value(config, key_index);
let key_hash = sha256_hex(&api_key_value);
let key_encrypted = secret_cipher.encrypt_plaintext(&api_key_value)?;
let existing = auth_reader
.find_export_standalone_api_key_by_id(&api_key_id)
.await?;
if existing
.as_ref()
.is_some_and(|record| record.key_hash != key_hash)
{
return Err(format!(
"existing pressure API key {api_key_id} has a different hash; use its original value or a new --api-key-id"
)
.into());
}
if existing.is_none() {
auth_writer
.create_standalone_api_key(CreateStandaloneApiKeyRecord {
user_id: operator_user_id.to_string(),
api_key_id: api_key_id.clone(),
key_hash: sha256_hex(&api_key_value),
key_encrypted: Some(api_key_value),
key_hash,
key_encrypted: Some(key_encrypted),
name: Some(format!("Local pressure API key {}", key_index + 1)),
allowed_providers: Some(vec![config.provider_id.clone()]),
allowed_api_formats: Some(vec!["openai:chat".to_string()]),
@@ -526,8 +571,8 @@ async fn seed_api_key(
.update_standalone_api_key_basic(
aether_data::repository::auth::UpdateStandaloneApiKeyBasicRecord {
api_key_id: api_key_id.clone(),
key_encrypted: None,
key_encrypted_present: false,
key_encrypted: Some(key_encrypted),
key_encrypted_present: true,
name: Some(format!("Local pressure API key {}", key_index + 1)),
name_present: true,
force_capabilities: None,
@@ -867,6 +912,8 @@ fn print_help() {
println!(
"Usage: cargo run -p aether-integration-tests --bin gateway_pressure_seed -- [options]\n\
\n\
The seed and gateway must share AETHER_GATEWAY_DATA_ENCRYPTION_KEY (or ENCRYPTION_KEY).\n\
\n\
Options:\n\
--database-url URL\n\
--output-env PATH\n\
@@ -104,8 +104,13 @@ struct AcceptanceReport {
reasons: Vec<String>,
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
run()
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run() -> Result<(), Box<dyn std::error::Error>> {
init_test_runtime_for("gateway-tunnel-stream-baseline");
let config = parse_args(std::env::args().skip(1).collect())?;
let report = run_suite(&config).await?;
@@ -318,8 +318,13 @@ struct ProtocolPeer {
stats: Arc<PeerStats>,
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
run()
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run() -> Result<(), Box<dyn std::error::Error>> {
init_test_runtime_for("llm-stream-stability-baseline");
let config = parse_args(std::env::args().skip(1).collect())?;
config.validate().map_err(std::io::Error::other)?;
@@ -441,6 +441,7 @@ async fn chat_completions(State(app): State<App>, request: axum::extract::Reques
completion
.take()
.expect("request completion guard should be present"),
false,
);
}
@@ -457,6 +458,14 @@ async fn chat_completions(State(app): State<App>, request: axum::extract::Reques
};
let stream = request_wants_stream(&body);
if stream {
let include_usage = serde_json::from_slice::<serde_json::Value>(&body)
.ok()
.and_then(|value| {
value
.pointer("/stream_options/include_usage")
.and_then(serde_json::Value::as_bool)
})
.unwrap_or(false);
record_response_header_created(&app, request_started.started_at.elapsed());
return build_chat_sse_response(
app,
@@ -464,6 +473,7 @@ async fn chat_completions(State(app): State<App>, request: axum::extract::Reques
completion
.take()
.expect("request completion guard should be present"),
include_usage,
);
}
// A stream truncation profile only applies after the request is known to be streaming.
@@ -609,6 +619,7 @@ fn build_chat_sse_response(
app: App,
profile: RequestProfile,
completion: RequestCompletionGuard,
include_usage: bool,
) -> Response {
let response_created_at = Instant::now();
let config = app.config.clone();
@@ -659,6 +670,21 @@ fn build_chat_sse_response(
yield Ok::<Bytes, std::io::Error>(Bytes::from(
"data: {\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n",
));
if include_usage {
let payload = json!({
"id": "chatcmpl-mock",
"object": "chat.completion.chunk",
"created": current_unix_secs(),
"model": "mock-model",
"choices": [],
"usage": {
"prompt_tokens": 1,
"completion_tokens": config.chunks.max(1),
"total_tokens": config.chunks.max(1) + 1
}
});
yield Ok::<Bytes, std::io::Error>(Bytes::from(format!("data: {payload}\n\n")));
}
yield Ok::<Bytes, std::io::Error>(Bytes::from("data: [DONE]\n\n"));
if let Some(completion) = completion.take() {
completion.complete();
@@ -1465,7 +1491,11 @@ mod tests {
) {
let response = client
.post(url)
.json(&json!({"stream": true, "model": "mock-test"}))
.json(&json!({
"stream": true,
"model": "mock-test",
"stream_options": {"include_usage": true}
}))
.send()
.await
.expect("headers should arrive before the body error");
@@ -1505,6 +1535,78 @@ mod tests {
!body.contains("[DONE]"),
"truncated stream must not emit [DONE]"
);
assert!(
!body.contains("\"usage\""),
"truncated stream must not emit terminal usage"
);
}
#[tokio::test]
async fn chat_stream_usage_is_opt_in_and_precedes_done() {
for chunks in [0, 3] {
for (include_usage, assume_stream) in [
(None, false),
(Some(false), false),
(Some(true), false),
(Some(true), true),
] {
let config = Config {
chunks,
chunk_delay: Duration::ZERO,
assume_stream,
..Default::default()
};
let app = App {
metrics: Arc::new(Metrics::for_binds(&config.binds)),
bind_label: Arc::from(config.binds[0].to_string()),
config,
};
let mut payload = json!({"stream": true, "model": "mock-test"});
if let Some(include_usage) = include_usage {
payload["stream_options"] = json!({"include_usage": include_usage});
}
let request = axum::http::Request::builder()
.body(Body::from(payload.to_string()))
.unwrap();
let response = chat_completions(State(app), request).await;
assert_eq!(response.status(), StatusCode::OK);
let body = to_bytes(response.into_body(), 16 * 1024).await.unwrap();
let body = std::str::from_utf8(&body).unwrap();
let frames = body
.split("\n\n")
.filter_map(|frame| frame.strip_prefix("data: "))
.collect::<Vec<_>>();
assert_eq!(frames.last(), Some(&"[DONE]"));
let payloads = frames[..frames.len() - 1]
.iter()
.map(|frame| serde_json::from_str::<serde_json::Value>(frame).unwrap())
.collect::<Vec<_>>();
let usage_chunks = payloads
.iter()
.filter(|payload| payload.get("usage").is_some())
.collect::<Vec<_>>();
if include_usage == Some(true) && !assume_stream {
assert_eq!(usage_chunks.len(), 1);
assert_eq!(payloads.last(), Some(usage_chunks[0]));
assert_eq!(usage_chunks[0]["object"], "chat.completion.chunk");
assert_eq!(usage_chunks[0]["choices"], json!([]));
assert_eq!(
usage_chunks[0]["usage"],
json!({
"prompt_tokens": 1,
"completion_tokens": chunks.max(1),
"total_tokens": chunks.max(1) + 1
})
);
assert_eq!(
payloads[payloads.len() - 2]["choices"][0]["finish_reason"],
"stop"
);
} else {
assert!(usage_chunks.is_empty());
}
}
}
}
async fn wait_for_completed(metrics: &Metrics, expected: u64) {
@@ -94,8 +94,13 @@ struct WebSocketAdmissionProbeResult {
runtime: BenchmarkRuntimeSnapshot,
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
run()
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run() -> Result<(), Box<dyn std::error::Error>> {
init_test_runtime_for("multi-instance-admission-baseline");
let config = parse_args(std::env::args().skip(1).collect())?;
let report = run_suite(&config).await?;
@@ -66,8 +66,13 @@ struct RelayOverheadSnapshot {
mean_delta_ms: i64,
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
run()
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run() -> Result<(), Box<dyn std::error::Error>> {
init_test_runtime_for("multi-instance-owner-relay-baseline");
let config = parse_args(std::env::args().skip(1).collect())?;
let report = run_suite(&config).await?;
@@ -54,8 +54,13 @@ struct SingleInstanceBaselineReport {
scenarios: Vec<NamedBaselineResult>,
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
run()
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run() -> Result<(), Box<dyn std::error::Error>> {
init_test_runtime_for("single-instance-baseline");
let config = parse_args(std::env::args().skip(1).collect())?;
let report = run_suite(&config).await?;
@@ -123,8 +123,13 @@ struct LockSample {
oldest_lock_wait_ms: i64,
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
run()
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run() -> Result<(), Box<dyn std::error::Error>> {
init_test_runtime_for("usage-aux-counter-hotspot-baseline");
let config = parse_args(std::env::args().skip(1).collect())?;
@@ -116,8 +116,13 @@ struct LockSample {
oldest_lock_wait_ms: i64,
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
run()
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run() -> Result<(), Box<dyn std::error::Error>> {
init_test_runtime_for("usage-counter-hotspot-baseline");
let config = parse_args(std::env::args().skip(1).collect())?;
@@ -453,6 +458,7 @@ fn usage_record(index: usize) -> UpsertUsageRecord {
let now_ms = now_unix_ms().saturating_add(index as u64);
let now_secs = now_ms / 1_000;
UpsertUsageRecord {
capture_retention: Default::default(),
request_id: format!("usage-hotspot-{index:08}"),
user_id: Some("user-hotspot".to_string()),
api_key_id: Some("api-key-hotspot".to_string()),
@@ -18,6 +18,8 @@ use sqlx::{PgPool, Row};
use tokio::sync::Mutex;
const PROVIDER_ID: &str = "provider-hotspot";
const USER_ID: &str = "settlement-hotspot-user";
const WALLET_ID: &str = "settlement-hotspot-wallet";
const REQUEST_PREFIX: &str = "settlement-hotspot";
const COST_PER_REQUEST_USD: f64 = 0.001;
@@ -99,6 +101,7 @@ struct CounterReport {
provider_monthly_outbox_rows: i64,
provider_monthly_used_usd: f64,
expected_provider_monthly_used_usd: f64,
wallet_consumed_usd: f64,
}
#[derive(Debug, Serialize, Clone, Copy, Default)]
@@ -120,8 +123,13 @@ struct LockSample {
oldest_lock_wait_ms: i64,
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
run()
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run() -> Result<(), Box<dyn std::error::Error>> {
init_test_runtime_for("usage-settlement-hotspot-baseline");
let config = parse_args(std::env::args().skip(1).collect())?;
@@ -239,6 +247,24 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
}
std::fs::write(path, format!("{raw}\n"))?;
}
if report.failed_requests != 0
|| report.counters.settled_usage_rows != config.requests as i64
|| report.counters.settlement_snapshot_rows != config.requests as i64
|| report.counters.outbox_pending_rows != 0
|| (report.counters.provider_monthly_used_usd
- report.counters.expected_provider_monthly_used_usd)
.abs()
> 1e-8
|| (report.counters.wallet_consumed_usd
- report.counters.expected_provider_monthly_used_usd)
.abs()
> 1e-8
{
return Err(std::io::Error::other(
"settlement baseline failed correctness checks; see report",
)
.into());
}
Ok(())
}
@@ -388,6 +414,24 @@ async fn wait_for_outbox_drain(
}
async fn seed_settlement_rows(pool: &PgPool, requests: usize) -> Result<(), sqlx::Error> {
sqlx::query(
"INSERT INTO users (id, username, email_verified) VALUES ($1, $1, true) ON CONFLICT (id) DO NOTHING",
)
.bind(USER_ID)
.execute(pool)
.await?;
sqlx::query(
r#"
INSERT INTO wallets (id, user_id, balance, gift_balance, total_consumed, created_at, updated_at)
VALUES ($1, $2, $3, 0, 0, NOW(), NOW())
ON CONFLICT (id) DO UPDATE SET balance = EXCLUDED.balance, gift_balance = 0, total_consumed = 0
"#,
)
.bind(WALLET_ID)
.bind(USER_ID)
.bind(requests as f64 * COST_PER_REQUEST_USD + 1.0)
.execute(pool)
.await?;
sqlx::query(
r#"
INSERT INTO providers (id, name, provider_type, monthly_used_usd)
@@ -432,6 +476,7 @@ WHERE request_id LIKE $1
INSERT INTO "usage" (
id,
request_id,
user_id,
provider_name,
model,
provider_id,
@@ -445,6 +490,7 @@ INSERT INTO "usage" (
SELECT
'settlement-usage-' || LPAD(gs::TEXT, 8, '0'),
$2 || '-' || LPAD(gs::TEXT, 8, '0'),
$7,
'Hotspot Provider',
'gpt-5',
$3,
@@ -463,6 +509,7 @@ FROM generate_series(0, $1::INTEGER - 1) AS gs
.bind(COST_PER_REQUEST_USD)
.bind(now_unix_ms() as i64)
.bind(now_unix_secs() as i64)
.bind(USER_ID)
.execute(pool)
.await?;
@@ -472,7 +519,7 @@ FROM generate_series(0, $1::INTEGER - 1) AS gs
fn settlement_input(index: usize) -> UsageSettlementInput {
UsageSettlementInput {
request_id: format!("{REQUEST_PREFIX}-{index:08}"),
user_id: None,
user_id: Some(USER_ID.to_string()),
api_key_id: None,
api_key_is_standalone: false,
provider_id: Some(PROVIDER_ID.to_string()),
@@ -552,11 +599,13 @@ SELECT
SELECT CAST(monthly_used_usd AS DOUBLE PRECISION)
FROM providers
WHERE id = $2
) AS provider_monthly_used_usd
) AS provider_monthly_used_usd,
(SELECT CAST(total_consumed AS DOUBLE PRECISION) FROM wallets WHERE id = $3) AS wallet_consumed_usd
"#,
)
.bind(format!("{REQUEST_PREFIX}-%"))
.bind(PROVIDER_ID)
.bind(WALLET_ID)
.fetch_one(pool)
.await?;
Ok(CounterReport {
@@ -568,6 +617,7 @@ SELECT
provider_monthly_outbox_rows: row.try_get("provider_monthly_outbox_rows")?,
provider_monthly_used_usd: row.try_get("provider_monthly_used_usd")?,
expected_provider_monthly_used_usd: (requests as f64) * COST_PER_REQUEST_USD,
wallet_consumed_usd: row.try_get("wallet_consumed_usd")?,
})
}
@@ -83,6 +83,9 @@ struct GatewayPressureReport {
settle_drain_elapsed_ms: u64,
settle_required_metrics_available: bool,
settle_missing_required_metrics: Vec<String>,
settle_baseline: SettleDrainBaseline,
settle_final_metrics: BTreeMap<String, u64>,
settle_observations: Vec<SettleDrainObservation>,
load: HttpLoadProbeResult,
metrics: GatewayPressureMetricsSummary,
}
@@ -573,9 +576,46 @@ struct SettleDrainResult {
elapsed: Duration,
required_metrics_available: bool,
missing_required_metrics: Vec<String>,
observations: Vec<SettleDrainObservation>,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
#[derive(Debug, Clone, Serialize)]
struct SettleDrainObservation {
elapsed_ms: u64,
drained: bool,
metrics: BTreeMap<String, u64>,
}
fn observe_settle_drain(
observations: &mut Vec<SettleDrainObservation>,
elapsed: Duration,
samples: &[PrometheusSample],
drained: bool,
) {
let observation = SettleDrainObservation {
elapsed_ms: elapsed.as_millis() as u64,
drained,
metrics: REQUIRED_SETTLE_DRAIN_METRICS
.iter()
.copied()
.chain([
"gateway_http_connections_in_flight",
"gateway_process_open_fds",
"usage_runtime_producers_in_flight",
"usage_runtime_delayed_lifecycle_pending",
])
.filter(|name| metric_is_available(samples, name))
.map(|name| (name.to_string(), metric_max(samples, name)))
.collect(),
};
if observations.len() < 256 {
observations.push(observation);
} else {
observations[255] = observation;
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize)]
struct SettleDrainBaseline {
tokio_alive_tasks: u64,
tokio_global_queue_depth: u64,
@@ -2302,6 +2342,8 @@ impl GatewayDbPoolPressureWindow {
fn samples_are_drained(samples: &[PrometheusSample], baseline: SettleDrainBaseline) -> bool {
missing_required_settle_drain_metrics(samples).is_empty()
&& metric_max(samples, "usage_runtime_producers_in_flight") == 0
&& metric_max(samples, "usage_runtime_delayed_lifecycle_pending") == 0
&& metric_max(samples, "request_candidate_queue_depth") == 0
&& metric_max(samples, "request_candidate_queue_pending_depth") == 0
&& metric_max(samples, "request_candidate_active_queue_depth") == 0
@@ -2403,6 +2445,7 @@ async fn wait_for_settle_drain(
) -> SettleDrainResult {
if settle_after.is_zero() {
return SettleDrainResult {
observations: Vec::new(),
completed: false,
elapsed: Duration::ZERO,
required_metrics_available: false,
@@ -2422,12 +2465,14 @@ async fn wait_for_settle_drain(
.min(Duration::from_millis(500))
.max(Duration::from_millis(50));
let mut stability = SettleDrainStability::default();
let mut observations = Vec::new();
loop {
match fetch_prometheus_samples(metrics_url).await {
Ok(samples) => {
missing_required_metrics = missing_required_settle_drain_metrics(&samples);
let drained = samples_are_drained(&samples, baseline);
observe_settle_drain(&mut observations, started.elapsed(), &samples, drained);
let terminal_loss_detected = {
let mut snapshot = summary.lock().await;
snapshot.observe(&samples);
@@ -2435,6 +2480,7 @@ async fn wait_for_settle_drain(
};
if terminal_loss_detected {
return SettleDrainResult {
observations,
completed: false,
elapsed: started.elapsed(),
required_metrics_available: missing_required_metrics.is_empty(),
@@ -2443,6 +2489,7 @@ async fn wait_for_settle_drain(
}
if stability.observe(started.elapsed(), drained) {
return SettleDrainResult {
observations,
completed: true,
elapsed: started.elapsed(),
required_metrics_available: true,
@@ -2465,6 +2512,12 @@ async fn wait_for_settle_drain(
let mut completed = false;
if let Ok(samples) = fetch_prometheus_samples(metrics_url).await {
observe_settle_drain(
&mut observations,
started.elapsed(),
&samples,
samples_are_drained(&samples, baseline),
);
missing_required_metrics = missing_required_settle_drain_metrics(&samples);
completed = stability.observe(started.elapsed(), samples_are_drained(&samples, baseline));
let mut snapshot = summary.lock().await;
@@ -2475,6 +2528,7 @@ async fn wait_for_settle_drain(
}
SettleDrainResult {
observations,
completed,
elapsed: started.elapsed(),
required_metrics_available: missing_required_metrics.is_empty(),
@@ -2540,13 +2594,25 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
suite: "gateway_pressure_probe",
acceptance_contract_version: ACCEPTANCE_CONTRACT_VERSION,
target_url: config.load.url,
metrics_url: config.metrics_url,
metrics_url: config.metrics_url.clone(),
sample_interval_ms: config.sample_interval.as_millis() as u64,
settle_after_ms: config.settle_after.as_millis() as u64,
settle_drain_completed: settle_drain.completed,
settle_drain_elapsed_ms: settle_drain.elapsed.as_millis() as u64,
settle_required_metrics_available: settle_drain.required_metrics_available,
settle_missing_required_metrics: settle_drain.missing_required_metrics,
settle_baseline: settle_drain_baseline,
settle_observations: settle_drain.observations,
settle_final_metrics: fetch_prometheus_samples(&config.metrics_url)
.await
.map(|samples| {
REQUIRED_SETTLE_DRAIN_METRICS
.iter()
.filter(|name| metric_is_available(&samples, name))
.map(|name| (name.to_string(), metric_max(&samples, name)))
.collect()
})
.unwrap_or_default(),
load,
metrics: Arc::try_unwrap(summary)
.unwrap_or_else(|_| panic!("metrics summary still referenced"))
@@ -52,8 +52,13 @@ struct RedisWorkerBaselineReport {
ack: OperationSummary,
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
run()
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run() -> Result<(), Box<dyn std::error::Error>> {
init_load_runtime_for("redis-worker-baseline");
let config = parse_args(std::env::args().skip(1).collect())?;
let report = run_suite(&config).await?;
@@ -121,8 +121,13 @@ impl SummaryCollector {
}
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
run()
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run() -> Result<(), Box<dyn std::error::Error>> {
init_load_runtime_for("runtime-redis-pressure");
let config = parse_args(std::env::args().skip(1).collect())?;
let report = run_suite(&config).await?;
+19 -2
View File
@@ -31,6 +31,7 @@ impl GatewayHarnessConfig {
#[derive(Debug)]
pub struct GatewayHarness {
server: SpawnedServer,
state: AppState,
}
impl GatewayHarness {
@@ -67,7 +68,7 @@ impl GatewayHarness {
if let Some(gate) = config.distributed_request_gate {
state = state.with_distributed_request_concurrency_gate(gate);
}
let router = build_router_with_state(state);
let router = build_router_with_state(state.clone());
let server = match port {
Some(port) => SpawnedServer::start_on_port(port, router)
.await
@@ -76,7 +77,7 @@ impl GatewayHarness {
.await
.map_err(|err| format!("failed to start gateway harness: {err}"))?,
};
Ok(Self { server })
Ok(Self { server, state })
}
pub fn base_url(&self) -> &str {
@@ -86,4 +87,20 @@ impl GatewayHarness {
pub fn port(&self) -> u16 {
self.server.port()
}
pub async fn metric_samples(&self) -> Result<Vec<crate::PrometheusSample>, String> {
let samples = aether_gateway::testkit::gateway_metric_samples(&self.state).await?;
Ok(samples
.into_iter()
.map(|sample| crate::PrometheusSample {
name: sample.name.to_string(),
labels: sample
.labels
.into_iter()
.map(|label| (label.key.to_string(), label.value))
.collect(),
value: sample.value.to_string(),
})
.collect())
}
}
+57 -6
View File
@@ -1,9 +1,6 @@
use std::time::Duration;
use aether_gateway::{
build_tunnel_runtime_router_with_state, TunnelConnConfig, TunnelControlPlaneClient,
TunnelRuntimeState,
};
use aether_gateway::{TunnelConnConfig, TunnelControlPlaneClient, TunnelRuntimeState};
use aether_runtime_state::RuntimeSemaphore;
use crate::server::SpawnedServer;
@@ -11,6 +8,8 @@ use crate::server::SpawnedServer;
pub const TUNNEL_HARNESS_NODE_ID: &str = "node-baseline";
pub const TUNNEL_HARNESS_GENERATION: &str = "tunnel-harness-generation-1";
pub const TUNNEL_HARNESS_MANAGEMENT_TOKEN: &str = "ae-tunnel-harness-management-token";
const RELAY_INSTANCE: &str = "tunnel-harness";
const RELAY_SECRET: &[u8] = b"tunnel-harness-relay-secret-32-bytes-minimum";
#[derive(Debug, Clone)]
pub struct TunnelHarnessConfig {
@@ -38,6 +37,7 @@ impl Default for TunnelHarnessConfig {
#[derive(Debug)]
pub struct TunnelHarness {
server: SpawnedServer,
node_id: String,
}
impl TunnelHarness {
@@ -74,7 +74,11 @@ impl TunnelHarness {
TUNNEL_HARNESS_GENERATION,
TUNNEL_HARNESS_MANAGEMENT_TOKEN,
)?;
let router = build_tunnel_runtime_router_with_state(state);
let router = aether_gateway::testkit::build_tunnel_pressure_router(
state,
RELAY_INSTANCE,
RELAY_SECRET,
)?;
let server = match port {
Some(port) => SpawnedServer::start_on_port(port, router)
.await
@@ -83,7 +87,10 @@ impl TunnelHarness {
.await
.map_err(|err| format!("failed to start tunnel harness: {err}"))?,
};
Ok(Self { server })
Ok(Self {
server,
node_id: config.node_id,
})
}
pub fn base_url(&self) -> &str {
@@ -93,6 +100,50 @@ impl TunnelHarness {
pub fn port(&self) -> u16 {
self.server.port()
}
pub fn relay_headers(
&self,
metadata_envelope: &[u8],
body: &[u8],
) -> std::collections::BTreeMap<String, String> {
use aether_contracts::tunnel::*;
use std::sync::atomic::{AtomicU64, Ordering};
static NONCE: AtomicU64 = AtomicU64::new(0);
let timestamp = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs();
let nonce = format!("harness-{}", NONCE.fetch_add(1, Ordering::Relaxed));
let digest = tunnel_relay_payload_digest(metadata_envelope, body);
let signature = sign_tunnel_relay_request(
RELAY_SECRET,
"load-probe",
RELAY_INSTANCE,
&self.node_id,
"",
false,
timestamp,
&nonce,
&digest,
);
[
(TUNNEL_RELAY_AUTH_SENDER_HEADER, "load-probe".to_string()),
(
TUNNEL_RELAY_OWNER_INSTANCE_HEADER,
RELAY_INSTANCE.to_string(),
),
(TUNNEL_RELAY_AUTH_TIMESTAMP_HEADER, timestamp.to_string()),
(TUNNEL_RELAY_AUTH_NONCE_HEADER, nonce),
(
TUNNEL_RELAY_AUTH_PAYLOAD_HEADER,
digest.encode_header_value(),
),
(TUNNEL_RELAY_AUTH_SIGNATURE_HEADER, signature),
]
.into_iter()
.map(|(key, value)| (key.to_string(), value))
.collect()
}
}
pub fn insert_tunnel_harness_auth_headers(
@@ -623,6 +623,10 @@ fn append_body_capture_metadata_entry(
);
}
pub(crate) fn mark_usage_event_capture_truncated(metadata: &mut Option<Value>, key: &str) {
aether_data_contracts::repository::usage::mark_usage_capture_memory_omitted(metadata, key);
}
fn upsert_body_capture_metadata_value_entry(
metadata: &mut Option<Value>,
key: &str,
+41
View File
@@ -15,6 +15,7 @@ pub struct UsageRuntimeConfig {
pub consumer_group: String,
pub dlq_stream_key: String,
pub stream_maxlen: usize,
pub queue_payload_max_bytes: usize,
pub consumer_batch_size: usize,
pub consumer_block_ms: u64,
pub reclaim_idle_ms: u64,
@@ -47,6 +48,7 @@ impl Default for UsageRuntimeConfig {
consumer_group: "usage_consumers".to_string(),
dlq_stream_key: "usage:events:dlq".to_string(),
stream_maxlen: 200_000,
queue_payload_max_bytes: 1024 * 1024,
consumer_batch_size: 128,
consumer_block_ms: 500,
reclaim_idle_ms: 60_000,
@@ -90,6 +92,11 @@ impl UsageRuntimeConfig {
"usage runtime dlq_stream_key cannot be empty".to_string(),
));
}
if self.stream_key == self.dlq_stream_key {
return Err(DataLayerError::InvalidConfiguration(
"usage runtime stream_key and dlq_stream_key must be different".to_string(),
));
}
if self.worker_count == 0 {
return Err(DataLayerError::InvalidConfiguration(
"usage runtime worker_count must be positive".to_string(),
@@ -124,6 +131,11 @@ impl UsageRuntimeConfig {
"usage runtime stream_maxlen must be positive".to_string(),
));
}
if self.queue_payload_max_bytes == 0 {
return Err(DataLayerError::InvalidConfiguration(
"usage runtime queue_payload_max_bytes must be positive".to_string(),
));
}
if self.consumer_batch_size == 0 {
return Err(DataLayerError::InvalidConfiguration(
"usage runtime consumer_batch_size must be positive".to_string(),
@@ -213,6 +225,19 @@ mod tests {
assert!(config.validate().is_err());
}
#[test]
fn enabled_config_rejects_dead_letter_stream_equal_to_source() {
let mut config = UsageRuntimeConfig::default();
config.dlq_stream_key = config.stream_key.clone();
assert!(config.validate().is_ok());
config.enabled = true;
assert!(matches!(
config.validate(),
Err(aether_data_contracts::DataLayerError::InvalidConfiguration(message))
if message.contains("must be different")
));
}
#[test]
fn enabled_config_rejects_zero_terminal_submission_limit() {
let config = UsageRuntimeConfig {
@@ -222,4 +247,20 @@ mod tests {
};
assert!(config.validate().is_err());
}
#[test]
fn queue_payload_limit_defaults_to_one_mib_and_rejects_zero_when_enabled() {
let mut config = UsageRuntimeConfig::default();
assert_eq!(config.queue_payload_max_bytes, 1024 * 1024);
config.queue_payload_max_bytes = 0;
assert!(config.validate().is_ok());
config.enabled = true;
assert!(matches!(
config.validate(),
Err(aether_data_contracts::DataLayerError::InvalidConfiguration(message))
if message.contains("queue_payload_max_bytes")
));
config.queue_payload_max_bytes = 1;
assert!(config.validate().is_ok());
}
}
@@ -0,0 +1,541 @@
use std::collections::BTreeMap;
use std::io::{self, Write};
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, LazyLock};
use aether_data_contracts::DataLayerError;
use aether_runtime_state::RuntimeQueueEntry;
use serde::Serialize;
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
const DEFAULT_ENCODING_BUDGET_BYTES: usize = 64 * 1024 * 1024;
const DEFAULT_ENCODING_JOBS: usize = 4;
const MAX_ENCODING_JOBS: usize = 128;
static ENCODING_BUDGET: LazyLock<Arc<DeadLetterEncodingBudget>> = LazyLock::new(|| {
Arc::new(DeadLetterEncodingBudget::new(
configured_limit(
std::env::var("AETHER_USAGE_DLQ_ENCODING_BUDGET_BYTES")
.ok()
.as_deref(),
DEFAULT_ENCODING_BUDGET_BYTES,
maximum_budget_bytes(),
),
configured_limit(
std::env::var("AETHER_USAGE_DLQ_ENCODING_MAX_JOBS")
.ok()
.as_deref(),
DEFAULT_ENCODING_JOBS,
MAX_ENCODING_JOBS,
),
))
});
pub(crate) fn shared_dead_letter_encoding_budget() -> Arc<DeadLetterEncodingBudget> {
Arc::clone(&ENCODING_BUDGET)
}
pub(crate) fn dead_letter_encoding_metrics() -> DeadLetterEncodingSnapshot {
ENCODING_BUDGET.snapshot()
}
fn maximum_budget_bytes() -> usize {
Semaphore::MAX_PERMITS.min(u32::MAX as usize)
}
fn configured_limit(raw: Option<&str>, fallback: usize, maximum: usize) -> usize {
raw.and_then(|raw| raw.trim().parse::<u128>().ok())
.filter(|value| *value > 0)
.map(|value| value.min(maximum as u128) as usize)
.unwrap_or(fallback.min(maximum))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct DeadLetterEncodingSnapshot {
pub(crate) limit_bytes: usize,
pub(crate) job_limit: usize,
pub(crate) reserved_bytes: usize,
pub(crate) active_jobs: usize,
pub(crate) capacity_rejected_total: u64,
pub(crate) oversized_rejected_total: u64,
pub(crate) encoded_total: u64,
}
/// Reserves raw string lengths plus their worst-case JSON encoding before any
/// clone or encoding allocation. This excludes collection/allocation overhead and
/// Redis command, packed-command, and connection buffers; it is not an RSS limit.
pub(crate) struct DeadLetterEncodingBudget {
limit_bytes: usize,
job_limit: usize,
bytes: Arc<Semaphore>,
jobs: Arc<Semaphore>,
reserved_bytes: AtomicUsize,
active_jobs: AtomicUsize,
capacity_rejected_total: AtomicU64,
oversized_rejected_total: AtomicU64,
encoded_total: AtomicU64,
#[cfg(test)]
encode_hook: std::sync::Mutex<Option<Box<dyn FnOnce() + Send>>>,
}
impl DeadLetterEncodingBudget {
pub(crate) fn new(limit_bytes: usize, job_limit: usize) -> Self {
let limit_bytes = limit_bytes.min(maximum_budget_bytes());
let job_limit = job_limit.clamp(1, MAX_ENCODING_JOBS);
Self {
limit_bytes,
job_limit,
bytes: Arc::new(Semaphore::new(limit_bytes)),
jobs: Arc::new(Semaphore::new(job_limit)),
reserved_bytes: AtomicUsize::new(0),
active_jobs: AtomicUsize::new(0),
capacity_rejected_total: AtomicU64::new(0),
oversized_rejected_total: AtomicU64::new(0),
encoded_total: AtomicU64::new(0),
#[cfg(test)]
encode_hook: std::sync::Mutex::new(None),
}
}
pub(crate) fn snapshot(&self) -> DeadLetterEncodingSnapshot {
DeadLetterEncodingSnapshot {
limit_bytes: self.limit_bytes,
job_limit: self.job_limit,
reserved_bytes: self.reserved_bytes.load(Ordering::Relaxed),
active_jobs: self.active_jobs.load(Ordering::Relaxed),
capacity_rejected_total: self.capacity_rejected_total.load(Ordering::Relaxed),
oversized_rejected_total: self.oversized_rejected_total.load(Ordering::Relaxed),
encoded_total: self.encoded_total.load(Ordering::Relaxed),
}
}
pub(crate) fn try_reserve(
self: &Arc<Self>,
entry: &RuntimeQueueEntry,
error: &str,
) -> Result<DeadLetterEncodingReservation, DataLayerError> {
let size = encoding_size(entry, error).filter(|size| size.total <= self.limit_bytes);
let Some(size) = size else {
self.oversized_rejected_total
.fetch_add(1, Ordering::Relaxed);
return Err(DataLayerError::InvalidInput(format!(
"dead-letter raw fields and worst-case JSON exceed the {}-byte encoding budget",
self.limit_bytes
)));
};
let job_permit = Arc::clone(&self.jobs)
.try_acquire_owned()
.map_err(|_| self.capacity_error())?;
let byte_permit = Arc::clone(&self.bytes)
.try_acquire_many_owned(size.total as u32)
.map_err(|_| self.capacity_error())?;
self.reserved_bytes.fetch_add(size.total, Ordering::Relaxed);
self.active_jobs.fetch_add(1, Ordering::Relaxed);
Ok(DeadLetterEncodingReservation {
budget: Arc::clone(self),
size,
_byte_permit: byte_permit,
_job_permit: job_permit,
})
}
fn capacity_error(&self) -> DataLayerError {
self.capacity_rejected_total.fetch_add(1, Ordering::Relaxed);
DataLayerError::TimedOut("dead-letter encoding capacity is exhausted".to_string())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct EncodingSize {
json: usize,
total: usize,
}
fn encoding_size(entry: &RuntimeQueueEntry, error: &str) -> Option<EncodingSize> {
let mut raw = entry.id.len().checked_add(error.len())?;
for (key, value) in &entry.fields {
raw = raw.checked_add(key.len())?.checked_add(value.len())?;
}
checked_encoding_size(raw, entry.fields.len())
}
fn checked_encoding_size(raw: usize, field_count: usize) -> Option<EncodingSize> {
// Every string byte needs at most six bytes (\u00XX). Each map entry adds
// four quotes, a colon and at most one comma to the empty envelope.
const EMPTY_ENVELOPE_BYTES: usize = br#"{"entry_id":"","fields":{},"error":""}"#.len();
let json = raw
.checked_mul(6)?
.checked_add(field_count.checked_mul(6)?)?
.checked_add(EMPTY_ENVELOPE_BYTES)?;
Some(EncodingSize {
json,
total: raw.checked_add(json)?,
})
}
pub(crate) struct DeadLetterEncodingReservation {
budget: Arc<DeadLetterEncodingBudget>,
size: EncodingSize,
_byte_permit: OwnedSemaphorePermit,
_job_permit: OwnedSemaphorePermit,
}
impl DeadLetterEncodingReservation {
pub(crate) fn encode_owned(
self,
entry: RuntimeQueueEntry,
error: String,
) -> impl std::future::Future<Output = Result<EncodedDeadLetter, DataLayerError>> + Send {
let input = EncodingInput {
entry,
error,
reservation: self,
};
async move {
tokio::task::spawn_blocking(move || input.encode())
.await
.map_err(|error| {
DataLayerError::UnexpectedValue(format!(
"dead-letter encoding task failed: {error}"
))
})?
}
}
}
impl Drop for DeadLetterEncodingReservation {
fn drop(&mut self) {
self.budget
.reserved_bytes
.fetch_sub(self.size.total, Ordering::Relaxed);
self.budget.active_jobs.fetch_sub(1, Ordering::Relaxed);
}
}
// Field order also protects cancellation/panic: raw data is dropped before its
// reservation, including when a queued blocking task never starts.
struct EncodingInput {
entry: RuntimeQueueEntry,
error: String,
reservation: DeadLetterEncodingReservation,
}
#[derive(Serialize)]
struct DeadLetterPayload<'a> {
entry_id: &'a str,
fields: &'a BTreeMap<String, String>,
error: &'a str,
}
impl EncodingInput {
fn encode(self) -> Result<EncodedDeadLetter, DataLayerError> {
#[cfg(test)]
{
let hook = self.reservation.budget.encode_hook.lock().unwrap().take();
if let Some(hook) = hook {
hook();
}
}
let mut writer = BoundedJsonWriter::new(self.reservation.size.json);
serde_json::to_writer(
&mut writer,
&DeadLetterPayload {
entry_id: &self.entry.id,
fields: &self.entry.fields,
error: &self.error,
},
)
.map_err(|error| {
DataLayerError::UnexpectedValue(format!(
"failed to encode complete dead-letter fields: {error}"
))
})?;
let payload = String::from_utf8(writer.bytes).map_err(|error| {
DataLayerError::UnexpectedValue(format!("dead-letter JSON was not UTF-8: {error}"))
})?;
self.reservation
.budget
.encoded_total
.fetch_add(1, Ordering::Relaxed);
Ok(EncodedDeadLetter {
entry_id: self.entry.id,
fields: BTreeMap::from([("payload".to_string(), payload)]),
_reservation: self.reservation,
})
}
}
pub(crate) struct EncodedDeadLetter {
pub(crate) entry_id: String,
pub(crate) fields: BTreeMap<String, String>,
// Remains owned by the result until transfer/append finishes.
_reservation: DeadLetterEncodingReservation,
}
struct BoundedJsonWriter {
bytes: Vec<u8>,
max_bytes: usize,
}
impl BoundedJsonWriter {
fn new(max_bytes: usize) -> Self {
Self {
bytes: Vec::new(),
max_bytes,
}
}
}
impl Write for BoundedJsonWriter {
fn write(&mut self, bytes: &[u8]) -> io::Result<usize> {
if bytes.len() > self.max_bytes.saturating_sub(self.bytes.len()) {
return Err(io::Error::other("dead-letter JSON encoding bound exceeded"));
}
let required = self.bytes.len() + bytes.len();
if required > self.bytes.capacity() {
let capacity = required
.max(self.bytes.capacity().saturating_mul(2))
.min(self.max_bytes);
self.bytes
.try_reserve_exact(capacity - self.bytes.len())
.map_err(|error| {
io::Error::other(format!("dead-letter JSON allocation failed: {error}"))
})?;
}
self.bytes.extend_from_slice(bytes);
Ok(bytes.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::*;
fn test_entry() -> RuntimeQueueEntry {
RuntimeQueueEntry {
id: "10-3".to_string(),
fields: BTreeMap::from([
("payload".to_string(), "historical payload".to_string()),
("extra".to_string(), "original metadata".to_string()),
]),
}
}
#[tokio::test]
async fn dead_letter_encoding_preserves_complete_wire_and_all_escape_forms() {
let control_bytes = (0u8..=31).map(char::from).collect::<String>();
let mut entry = test_entry();
entry.id.push_str("\"\\\n");
entry.fields.insert(
format!("{control_bytes}\"\\"),
format!("{control_bytes}\"\\\u{4e2d}\u{6587}\u{1f600}"),
);
let error = format!("error:{control_bytes}\"\\\u{00e9}");
let expected = serde_json::to_string(&DeadLetterPayload {
entry_id: &entry.id,
fields: &entry.fields,
error: &error,
})
.unwrap();
let size = encoding_size(&entry, &error).unwrap();
assert!(expected.len() <= size.json);
let budget = Arc::new(DeadLetterEncodingBudget::new(size.total, 1));
let encoded = budget
.try_reserve(&entry, &error)
.unwrap()
.encode_owned(entry.clone(), error.clone())
.await
.unwrap();
assert_eq!(encoded.entry_id, entry.id);
assert_eq!(encoded.fields.len(), 1);
assert_eq!(encoded.fields["payload"], expected);
let decoded: serde_json::Value = serde_json::from_str(&encoded.fields["payload"]).unwrap();
assert_eq!(
decoded["fields"],
serde_json::to_value(entry.fields).unwrap()
);
assert_eq!(decoded["error"], error);
assert_eq!(budget.snapshot().encoded_total, 1);
assert_eq!(budget.snapshot().reserved_bytes, size.total);
assert_eq!(budget.snapshot().active_jobs, 1);
drop(encoded);
assert_eq!(budget.snapshot().reserved_bytes, 0);
assert_eq!(budget.snapshot().active_jobs, 0);
}
#[test]
fn dead_letter_encoding_budget_rejects_oversize_and_saturation_without_waiters() {
let entry = test_entry();
let size = encoding_size(&entry, "failure").unwrap();
let too_small = Arc::new(DeadLetterEncodingBudget::new(size.total - 1, 1));
assert!(matches!(
too_small.try_reserve(&entry, "failure"),
Err(DataLayerError::InvalidInput(_))
));
assert_eq!(too_small.snapshot().oversized_rejected_total, 1);
assert_eq!(too_small.snapshot().reserved_bytes, 0);
assert_eq!(too_small.snapshot().active_jobs, 0);
for (limit, jobs) in [(size.total, 2), (size.total * 2, 1)] {
let budget = Arc::new(DeadLetterEncodingBudget::new(limit, jobs));
let first = budget.try_reserve(&entry, "failure").unwrap();
assert!(matches!(
budget.try_reserve(&entry, "failure"),
Err(DataLayerError::TimedOut(_))
));
assert_eq!(budget.snapshot().capacity_rejected_total, 1);
assert_eq!(budget.snapshot().reserved_bytes, size.total);
assert_eq!(budget.snapshot().active_jobs, 1);
assert_eq!(budget.jobs.available_permits(), jobs - 1);
drop(first);
assert_eq!(budget.snapshot().reserved_bytes, 0);
assert_eq!(budget.jobs.available_permits(), jobs);
assert_eq!(budget.bytes.available_permits(), limit);
drop(budget.try_reserve(&entry, "failure").unwrap());
}
}
#[test]
fn dead_letter_encoding_bounds_and_environment_cannot_overflow() {
assert!(checked_encoding_size(usize::MAX, 0).is_none());
assert!(checked_encoding_size(usize::MAX / 6, 0).is_none());
assert!(checked_encoding_size(0, usize::MAX).is_none());
assert!(checked_encoding_size(usize::MAX / 7, 1).is_none());
assert!(checked_encoding_size(0, 0).unwrap().total > 0);
assert_eq!(configured_limit(None, 4, 128), 4);
assert_eq!(configured_limit(Some("0"), 4, 128), 4);
assert_eq!(configured_limit(Some("invalid"), 4, 128), 4);
assert_eq!(configured_limit(Some(" 2 "), 4, 128), 2);
assert_eq!(configured_limit(Some("99999999999"), 4, 128), 128);
assert_eq!(
configured_limit(Some(&u128::MAX.to_string()), 4, maximum_budget_bytes()),
maximum_budget_bytes()
);
let budget = DeadLetterEncodingBudget::new(usize::MAX, usize::MAX);
assert_eq!(budget.snapshot().limit_bytes, maximum_budget_bytes());
assert_eq!(budget.snapshot().job_limit, MAX_ENCODING_JOBS);
}
#[test]
fn dead_letter_encoding_bounded_writer_grows_geometrically_and_stops_at_limit() {
let mut writer = BoundedJsonWriter::new(4096);
let mut allocations = 0;
for _ in 0..4096 {
let previous_capacity = writer.bytes.capacity();
writer.write_all(b"x").unwrap();
allocations += usize::from(previous_capacity != writer.bytes.capacity());
}
assert!(allocations <= 13, "allocations: {allocations}");
assert!(writer.write_all(b"y").is_err());
assert_eq!(writer.bytes.len(), 4096);
assert!(writer.bytes.iter().all(|byte| *byte == b'x'));
}
#[tokio::test]
async fn dead_letter_encoding_unpolled_cancellation_releases_reservation() {
let entry = test_entry();
let budget = Arc::new(DeadLetterEncodingBudget::new(4096, 1));
let reservation = budget.try_reserve(&entry, "failure").unwrap();
let future = reservation.encode_owned(entry, "failure".to_string());
assert_eq!(budget.snapshot().active_jobs, 1);
drop(future);
assert_eq!(budget.snapshot().active_jobs, 0);
assert_eq!(budget.snapshot().reserved_bytes, 0);
assert_eq!(budget.snapshot().encoded_total, 0);
}
#[tokio::test]
async fn dead_letter_encoding_running_cancellation_holds_budget_until_closure_exits() {
let entry = test_entry();
let budget = Arc::new(DeadLetterEncodingBudget::new(4096, 1));
let (started_tx, started_rx) = tokio::sync::oneshot::channel();
let (release_tx, release_rx) = std::sync::mpsc::channel();
*budget.encode_hook.lock().unwrap() = Some(Box::new(move || {
let _ = started_tx.send(());
release_rx.recv_timeout(Duration::from_secs(2)).unwrap();
}));
let reservation = budget.try_reserve(&entry, "failure").unwrap();
let task = tokio::spawn(reservation.encode_owned(entry, "failure".to_string()));
tokio::time::timeout(Duration::from_secs(2), started_rx)
.await
.unwrap()
.unwrap();
task.abort();
assert!(matches!(task.await, Err(error) if error.is_cancelled()));
assert_eq!(budget.snapshot().active_jobs, 1);
assert!(budget.snapshot().reserved_bytes > 0);
assert!(budget.try_reserve(&test_entry(), "failure").is_err());
release_tx.send(()).unwrap();
tokio::time::timeout(Duration::from_secs(2), async {
while budget.snapshot().active_jobs != 0 {
tokio::task::yield_now().await;
}
})
.await
.unwrap();
assert_eq!(budget.snapshot().reserved_bytes, 0);
assert_eq!(budget.snapshot().encoded_total, 1);
}
#[tokio::test]
async fn dead_letter_encoding_panic_releases_input_and_reservation() {
let entry = test_entry();
let budget = Arc::new(DeadLetterEncodingBudget::new(4096, 1));
*budget.encode_hook.lock().unwrap() = Some(Box::new(|| panic!("encoding test panic")));
let result = budget
.try_reserve(&entry, "failure")
.unwrap()
.encode_owned(entry, "failure".to_string())
.await;
assert!(matches!(result, Err(DataLayerError::UnexpectedValue(_))));
assert_eq!(budget.snapshot().reserved_bytes, 0);
assert_eq!(budget.snapshot().active_jobs, 0);
assert_eq!(budget.snapshot().encoded_total, 0);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn dead_letter_encoding_concurrent_reservations_have_no_hidden_waiting_queue() {
const TASKS: usize = 16;
let entry = Arc::new(test_entry());
let size = encoding_size(&entry, "failure").unwrap();
let budget = Arc::new(DeadLetterEncodingBudget::new(size.total * 2, 2));
let ready = Arc::new(tokio::sync::Barrier::new(TASKS + 1));
let release = Arc::new(tokio::sync::Barrier::new(TASKS + 1));
let mut tasks = tokio::task::JoinSet::new();
for _ in 0..TASKS {
let entry = Arc::clone(&entry);
let budget = Arc::clone(&budget);
let ready = Arc::clone(&ready);
let release = Arc::clone(&release);
tasks.spawn(async move {
let reservation = budget.try_reserve(&entry, "failure");
ready.wait().await;
release.wait().await;
reservation.is_ok()
});
}
tokio::time::timeout(Duration::from_secs(2), ready.wait())
.await
.unwrap();
assert_eq!(budget.snapshot().active_jobs, 2);
assert_eq!(budget.snapshot().reserved_bytes, size.total * 2);
assert_eq!(
budget.snapshot().capacity_rejected_total,
(TASKS - 2) as u64
);
release.wait().await;
let mut admitted = 0;
while let Some(result) = tasks.join_next().await {
admitted += usize::from(result.unwrap());
}
assert_eq!(admitted, 2);
assert_eq!(budget.snapshot().reserved_bytes, 0);
assert_eq!(budget.snapshot().active_jobs, 0);
}
}
+766 -6
View File
@@ -1,4 +1,5 @@
use std::collections::BTreeMap;
use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
use aether_data_contracts::repository::usage::UsageBodyCaptureState;
@@ -6,8 +7,18 @@ use aether_data_contracts::DataLayerError;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::body_capture::mark_usage_event_capture_truncated;
pub use crate::event_capture_budget::UsageEventCaptureRetention;
use crate::event_capture_budget::{
json_heap_estimate, shared_capture_memory_budget, EventCaptureMemoryBudget,
};
pub const USAGE_EVENT_VERSION: u8 = 1;
#[path = "event_wire.rs"]
mod wire;
pub(crate) use wire::EncodedUsageEvent;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum UsageEventType {
@@ -18,7 +29,7 @@ pub enum UsageEventType {
Cancelled,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Default)]
#[derive(Debug, PartialEq, Serialize, Deserialize, Default)]
pub struct UsageEventData {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub user_id: Option<String>,
@@ -144,6 +155,171 @@ pub struct UsageEventData {
pub local_execution_runtime_miss_reason: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub request_metadata: Option<Value>,
#[doc(hidden)]
#[serde(skip)]
pub capture_retention: UsageEventCaptureRetention,
}
impl UsageEventData {
fn capture_heap_estimate(&self) -> usize {
[
&self.request_body,
&self.provider_request_body,
&self.response_body,
&self.client_response_body,
]
.into_iter()
.flatten()
.fold(0usize, |bytes, body| {
bytes
.saturating_add(std::mem::size_of::<Value>())
.saturating_add(json_heap_estimate(body))
})
}
fn captured_fields(&self) -> [bool; 4] {
[
self.request_body.is_some(),
self.provider_request_body.is_some(),
self.response_body.is_some(),
self.client_response_body.is_some(),
]
}
fn mark_capture_omitted(&mut self, captured: [bool; 4]) {
for (present, key, state) in [
(captured[0], "request", &mut self.request_body_state),
(
captured[1],
"provider_request",
&mut self.provider_request_body_state,
),
(captured[2], "response", &mut self.response_body_state),
(
captured[3],
"client_response",
&mut self.client_response_body_state,
),
] {
if present
&& !matches!(
*state,
Some(
UsageBodyCaptureState::None
| UsageBodyCaptureState::Disabled
| UsageBodyCaptureState::Unavailable
)
)
{
*state = Some(UsageBodyCaptureState::Truncated);
mark_usage_event_capture_truncated(&mut self.request_metadata, key);
}
}
}
pub(crate) fn apply_capture_memory_budget(
&mut self,
budget: std::sync::Arc<EventCaptureMemoryBudget>,
) {
let bytes = self.capture_heap_estimate();
if self
.capture_retention
.reserve(std::sync::Arc::clone(&budget), bytes)
{
return;
}
let captured = self.captured_fields();
self.request_body = None;
self.provider_request_body = None;
self.response_body = None;
self.client_response_body = None;
self.mark_capture_omitted(captured);
// The previous lease is released only after the owned JSON bodies are gone.
self.capture_retention.clear(budget);
}
}
impl Clone for UsageEventData {
fn clone(&self) -> Self {
let (capture_retention, retain_bodies) = self
.capture_retention
.clone_for_bodies(|| self.capture_heap_estimate());
// Enumerate every field so additions require an explicit ownership decision.
let mut cloned = Self {
user_id: self.user_id.clone(),
api_key_id: self.api_key_id.clone(),
username: self.username.clone(),
api_key_name: self.api_key_name.clone(),
provider_name: self.provider_name.clone(),
model: self.model.clone(),
target_model: self.target_model.clone(),
model_id: self.model_id.clone(),
global_model_id: self.global_model_id.clone(),
provider_id: self.provider_id.clone(),
provider_endpoint_id: self.provider_endpoint_id.clone(),
provider_api_key_id: self.provider_api_key_id.clone(),
request_type: self.request_type.clone(),
api_format: self.api_format.clone(),
api_family: self.api_family.clone(),
endpoint_kind: self.endpoint_kind.clone(),
endpoint_api_format: self.endpoint_api_format.clone(),
provider_api_family: self.provider_api_family.clone(),
provider_endpoint_kind: self.provider_endpoint_kind.clone(),
has_format_conversion: self.has_format_conversion,
is_stream: self.is_stream,
input_tokens: self.input_tokens,
output_tokens: self.output_tokens,
total_tokens: self.total_tokens,
cache_creation_input_tokens: self.cache_creation_input_tokens,
cache_creation_ephemeral_5m_input_tokens: self.cache_creation_ephemeral_5m_input_tokens,
cache_creation_ephemeral_1h_input_tokens: self.cache_creation_ephemeral_1h_input_tokens,
cache_read_input_tokens: self.cache_read_input_tokens,
cache_creation_cost_usd: self.cache_creation_cost_usd,
cache_read_cost_usd: self.cache_read_cost_usd,
output_price_per_1m: self.output_price_per_1m,
total_cost_usd: self.total_cost_usd,
actual_total_cost_usd: self.actual_total_cost_usd,
status_code: self.status_code,
error_message: self.error_message.clone(),
error_category: self.error_category.clone(),
response_time_ms: self.response_time_ms,
first_byte_time_ms: self.first_byte_time_ms,
request_headers: self.request_headers.clone(),
request_body: retain_bodies.then(|| self.request_body.clone()).flatten(),
request_body_ref: self.request_body_ref.clone(),
request_body_state: self.request_body_state,
provider_request_headers: self.provider_request_headers.clone(),
provider_request_body: retain_bodies
.then(|| self.provider_request_body.clone())
.flatten(),
provider_request_body_ref: self.provider_request_body_ref.clone(),
provider_request_body_state: self.provider_request_body_state,
response_headers: self.response_headers.clone(),
response_body: retain_bodies.then(|| self.response_body.clone()).flatten(),
response_body_ref: self.response_body_ref.clone(),
response_body_state: self.response_body_state,
client_response_headers: self.client_response_headers.clone(),
client_response_body: retain_bodies
.then(|| self.client_response_body.clone())
.flatten(),
client_response_body_ref: self.client_response_body_ref.clone(),
client_response_body_state: self.client_response_body_state,
candidate_id: self.candidate_id.clone(),
candidate_index: self.candidate_index,
key_name: self.key_name.clone(),
planner_kind: self.planner_kind.clone(),
route_family: self.route_family.clone(),
route_kind: self.route_kind.clone(),
execution_path: self.execution_path.clone(),
local_execution_runtime_miss_reason: self.local_execution_runtime_miss_reason.clone(),
request_metadata: self.request_metadata.clone(),
capture_retention,
};
if !retain_bodies {
cloned.mark_capture_omitted(self.captured_fields());
}
cloned
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
@@ -164,6 +340,16 @@ struct UsageEventEnvelope {
data: UsageEventData,
}
#[derive(Serialize)]
struct BorrowedUsageEventEnvelope<'a, T: ?Sized> {
v: u8,
#[serde(rename = "type")]
event_type: UsageEventType,
request_id: &'a str,
timestamp_ms: u64,
data: &'a T,
}
impl UsageEvent {
pub fn new(
event_type: UsageEventType,
@@ -179,12 +365,12 @@ impl UsageEvent {
}
pub fn to_stream_fields(&self) -> Result<BTreeMap<String, String>, DataLayerError> {
let payload = UsageEventEnvelope {
let payload = BorrowedUsageEventEnvelope {
v: USAGE_EVENT_VERSION,
event_type: self.event_type,
request_id: self.request_id.clone(),
request_id: &self.request_id,
timestamp_ms: self.timestamp_ms,
data: self.data.clone(),
data: &self.data,
};
let payload = serde_json::to_string(&payload).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
@@ -194,7 +380,21 @@ impl UsageEvent {
Ok(BTreeMap::from([("payload".to_string(), payload)]))
}
pub(crate) fn to_bounded_stream_fields(
&self,
max_bytes: usize,
) -> Result<EncodedUsageEvent, DataLayerError> {
wire::encode(self, max_bytes)
}
pub fn from_stream_fields(fields: &BTreeMap<String, String>) -> Result<Self, DataLayerError> {
Self::from_stream_fields_with_capture_budget(fields, shared_capture_memory_budget())
}
pub(crate) fn from_stream_fields_with_capture_budget(
fields: &BTreeMap<String, String>,
budget: Arc<EventCaptureMemoryBudget>,
) -> Result<Self, DataLayerError> {
let payload = fields.get("payload").ok_or_else(|| {
DataLayerError::UnexpectedValue(
"usage event stream entry missing payload field".to_string(),
@@ -212,12 +412,16 @@ impl UsageEvent {
)));
}
Ok(Self {
let mut event = Self {
event_type: envelope.event_type,
request_id: envelope.request_id,
timestamp_ms: envelope.timestamp_ms,
data: envelope.data,
})
};
// The wire format has no ownership lease. Preserve billing facts before a decoded
// body can be omitted; the raw Redis response and serde allocation are not budgeted here.
crate::runtime::prepare_decoded_event_capture_memory(&mut event, budget);
Ok(event)
}
}
@@ -230,8 +434,564 @@ pub fn now_ms() -> u64 {
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use std::sync::Arc;
use aether_data_contracts::repository::usage::UsageBodyCaptureState;
use aether_data_contracts::DataLayerError;
use serde_json::json;
use crate::event_capture_budget::EventCaptureMemoryBudget;
use crate::{
apply_usage_body_capture_policy_to_event, build_upsert_usage_record_from_event,
UsageBodyCapturePolicy,
};
use super::{UsageEvent, UsageEventData, UsageEventType};
fn captured_event() -> UsageEvent {
UsageEvent {
event_type: UsageEventType::Failed,
request_id: "capture-budget-request".to_string(),
timestamp_ms: 123_456,
data: UsageEventData {
provider_name: "provider".to_string(),
model: "model".to_string(),
input_tokens: Some(100),
output_tokens: Some(500),
total_tokens: Some(600),
cache_read_input_tokens: Some(0),
cache_creation_input_tokens: Some(25),
actual_total_cost_usd: Some(1.25),
status_code: Some(502),
error_category: Some("upstream_error".to_string()),
error_message: Some("upstream failed".to_string()),
request_body: Some(json!({"messages": [{"content": "request"}]})),
provider_request_body: Some(json!({"input": "upstream request"})),
response_body: Some(json!({"usage": {"input_tokens": 100, "output_tokens": 500}})),
client_response_body: Some(json!({"error": "client response"})),
request_body_state: Some(UsageBodyCaptureState::Inline),
provider_request_body_state: Some(UsageBodyCaptureState::Inline),
response_body_state: Some(UsageBodyCaptureState::Inline),
client_response_body_state: Some(UsageBodyCaptureState::Inline),
request_metadata: Some(json!({
"requested_reasoning_effort": "high",
"provider_reasoning_effort": "medium",
"provider_service_tier": "priority",
"provider_actual_service_tier": "default",
"provider_cache_ttl_minutes": 60,
"plan_usage_reservation_token": "550e8400-e29b-41d4-a716-446655440000",
"body_capture": {"response": {"state": "inline", "source_bytes": 1000}}
})),
..UsageEventData::default()
},
}
}
#[test]
fn event_capture_budget_zero_preserves_billing_refs_and_database_truncation() {
let budget = Arc::new(EventCaptureMemoryBudget::new(0));
let mut event = captured_event();
event.data.request_body_ref =
Some("usage://capture-budget-request/request_body".to_string());
event.data.response_body_ref =
Some("usage://capture-budget-request/response_body".to_string());
event.data.apply_capture_memory_budget(Arc::clone(&budget));
assert!(event.data.request_body.is_none());
assert!(event.data.provider_request_body.is_none());
assert!(event.data.response_body.is_none());
assert!(event.data.client_response_body.is_none());
let capture_metadata = event
.data
.request_metadata
.as_ref()
.expect("capture metadata");
assert_eq!(
capture_metadata["body_capture"]["response"]["source_bytes"],
1000
);
assert_eq!(
capture_metadata["body_capture"]["response"]["stored_bytes"],
0
);
assert_eq!(
capture_metadata["body_capture"]["response"]["reason"],
"usage_event_memory_budget_exceeded"
);
let record = build_upsert_usage_record_from_event(&event).expect("record mapping");
assert_eq!(record.status, "failed");
assert_eq!(record.input_tokens, Some(100));
assert_eq!(record.output_tokens, Some(500));
assert_eq!(record.cache_read_input_tokens, Some(0));
assert_eq!(record.cache_creation_input_tokens, Some(25));
assert_eq!(record.actual_total_cost_usd, Some(1.25));
assert_eq!(record.error_category.as_deref(), Some("upstream_error"));
assert_eq!(record.request_body_ref, event.data.request_body_ref);
assert_eq!(record.response_body_ref, event.data.response_body_ref);
for state in [
record.request_body_state,
record.provider_request_body_state,
record.response_body_state,
record.client_response_body_state,
] {
assert_eq!(state, Some(UsageBodyCaptureState::Truncated));
}
let metadata = record.request_metadata.expect("preserved metadata");
assert_eq!(metadata["provider_service_tier"], "priority");
assert_eq!(metadata["provider_actual_service_tier"], "default");
assert_eq!(metadata["provider_cache_ttl_minutes"], 60);
assert_eq!(
metadata["plan_usage_reservation_token"],
"550e8400-e29b-41d4-a716-446655440000"
);
// Persistence projects billing metadata; capture state remains in typed columns.
assert!(metadata.get("body_capture").is_none());
assert_eq!(budget.retained_bytes(), 0);
assert_eq!(budget.downgraded_total(), 1);
}
#[test]
fn event_capture_budget_clone_reserves_each_copy_before_cloning_bodies() {
let mut event = captured_event();
let weight = event.data.capture_heap_estimate();
let budget = Arc::new(EventCaptureMemoryBudget::new(weight * 2));
event.data.apply_capture_memory_budget(Arc::clone(&budget));
let copy = event.clone();
assert_eq!(copy, event);
assert_eq!(budget.retained_bytes(), weight * 2);
let downgraded = event.clone();
assert!(event.data.response_body.is_some());
assert!(copy.data.response_body.is_some());
assert!(downgraded.data.response_body.is_none());
assert_eq!(
downgraded.data.response_body_state,
Some(UsageBodyCaptureState::Truncated)
);
assert_eq!(downgraded.data.total_tokens, event.data.total_tokens);
assert_eq!(downgraded.timestamp_ms, event.timestamp_ms);
assert_eq!(downgraded.event_type, event.event_type);
assert_eq!(budget.retained_bytes(), weight * 2);
drop(copy);
assert_eq!(budget.retained_bytes(), weight);
drop((event, downgraded));
assert_eq!(budget.retained_bytes(), 0);
}
#[test]
fn event_capture_budget_serialization_borrows_bodies_without_charging_a_clone() {
let mut event = captured_event();
let weight = event.data.capture_heap_estimate();
let budget = Arc::new(EventCaptureMemoryBudget::new(weight));
crate::runtime::prepare_event_capture_memory(&mut event, Arc::clone(&budget));
let decoded_budget = Arc::new(EventCaptureMemoryBudget::new(usize::MAX));
for _ in 0..3 {
let fields = event.to_stream_fields().expect("wire serialization");
assert!(!fields["payload"].contains("capture_retention"));
let decoded = UsageEvent::from_stream_fields_with_capture_budget(
&fields,
Arc::clone(&decoded_budget),
)
.expect("wire decode");
assert_eq!(decoded, event);
assert_eq!(budget.retained_bytes(), weight);
assert!(decoded_budget.retained_bytes() > 0);
drop(decoded);
assert_eq!(decoded_budget.retained_bytes(), 0);
}
assert_eq!(budget.downgraded_total(), 0);
drop(event);
assert_eq!(budget.retained_bytes(), 0);
}
#[test]
fn from_stream_fields_legacy_body_budget_preserves_billing_and_request_facts() {
let mut event = captured_event();
event.data.model = "gpt-5.6-sol".to_string();
event.data.endpoint_api_format = Some("openai:responses".to_string());
event.data.request_body = Some(json!({"reasoning": {"effort": "high"}}));
event.data.provider_request_body = Some(json!({
"model": "gpt-5.6-sol", "reasoning": {"effort": "medium"},
"service_tier": "priority"
}));
event.data.response_body = Some(json!({"service_tier": "Default"}));
event.data.request_body_state = None;
event.data.provider_request_body_state = None;
event.data.response_body_state = None;
event.data.client_response_body_state = None;
event.data.cache_creation_ephemeral_5m_input_tokens = Some(0);
event.data.cache_creation_ephemeral_1h_input_tokens = Some(25);
event.data.cache_read_cost_usd = Some(0.0);
event.data.request_body_ref = Some("usage://legacy/request".to_string());
event.data.request_metadata = Some(json!({
"plan_usage_reservation_token": "550e8400-e29b-41d4-a716-446655440000"
}));
let fields = event.to_stream_fields().expect("legacy wire serialization");
assert!(!fields["payload"].contains("request_body_state"));
let budget = Arc::new(EventCaptureMemoryBudget::new(0));
let decoded =
UsageEvent::from_stream_fields_with_capture_budget(&fields, Arc::clone(&budget))
.expect("legacy wire decode");
assert_eq!(decoded.event_type, UsageEventType::Failed);
assert_eq!(decoded.request_id, event.request_id);
assert_eq!(decoded.timestamp_ms, event.timestamp_ms);
assert_eq!(decoded.data.input_tokens, Some(100));
assert_eq!(decoded.data.output_tokens, Some(500));
assert_eq!(decoded.data.total_tokens, Some(600));
assert_eq!(decoded.data.cache_creation_input_tokens, Some(25));
assert_eq!(
decoded.data.cache_creation_ephemeral_5m_input_tokens,
Some(0)
);
assert_eq!(
decoded.data.cache_creation_ephemeral_1h_input_tokens,
Some(25)
);
assert_eq!(decoded.data.cache_read_input_tokens, Some(0));
assert_eq!(decoded.data.cache_read_cost_usd, Some(0.0));
assert_eq!(decoded.data.actual_total_cost_usd, Some(1.25));
assert_eq!(decoded.data.status_code, Some(502));
assert_eq!(
decoded.data.error_category.as_deref(),
Some("upstream_error")
);
assert_eq!(decoded.data.request_body_ref, event.data.request_body_ref);
assert!(decoded.data.request_body.is_none());
assert!(decoded.data.provider_request_body.is_none());
assert!(decoded.data.response_body.is_none());
assert!(decoded.data.client_response_body.is_none());
for state in [
decoded.data.request_body_state,
decoded.data.provider_request_body_state,
decoded.data.response_body_state,
decoded.data.client_response_body_state,
] {
assert_eq!(state, Some(UsageBodyCaptureState::Truncated));
}
let metadata = decoded
.data
.request_metadata
.as_ref()
.expect("preserved facts");
assert_eq!(metadata["requested_reasoning_effort"], "high");
assert_eq!(metadata["provider_reasoning_effort"], "medium");
assert_eq!(metadata["provider_service_tier"], "priority");
assert_eq!(metadata["provider_actual_service_tier"], "default");
assert_eq!(metadata["provider_cache_ttl_minutes"], 30);
assert_eq!(
metadata["plan_usage_reservation_token"],
"550e8400-e29b-41d4-a716-446655440000"
);
assert_eq!(budget.retained_bytes(), 0);
assert_eq!(budget.downgraded_total(), 1);
}
#[test]
fn from_stream_fields_reconstructed_lease_also_bounds_recorder_clones() {
let fields = captured_event()
.to_stream_fields()
.expect("wire serialization");
let probe_budget = Arc::new(EventCaptureMemoryBudget::new(usize::MAX));
let probe =
UsageEvent::from_stream_fields_with_capture_budget(&fields, Arc::clone(&probe_budget))
.expect("estimate decoded allocation");
let weight = probe_budget.retained_bytes();
assert!(weight > 0);
drop(probe);
assert_eq!(probe_budget.retained_bytes(), 0);
let budget = Arc::new(EventCaptureMemoryBudget::new(weight * 2));
let event =
UsageEvent::from_stream_fields_with_capture_budget(&fields, Arc::clone(&budget))
.expect("wire decode");
let recorder_copy = event.clone();
assert!(recorder_copy.data.response_body.is_some());
assert_eq!(budget.retained_bytes(), weight * 2);
let omitted_copy = event.clone();
assert!(omitted_copy.data.response_body.is_none());
assert_eq!(omitted_copy.data.total_tokens, Some(600));
assert_eq!(omitted_copy.data.cache_read_input_tokens, Some(0));
assert_eq!(budget.retained_bytes(), weight * 2);
drop(event);
assert_eq!(budget.retained_bytes(), weight);
drop((recorder_copy, omitted_copy));
assert_eq!(budget.retained_bytes(), 0);
}
fn assert_typed_body_clear_is_preserved(event: &UsageEvent, state: UsageBodyCaptureState) {
for (body, reference, actual_state) in [
(
&event.data.request_body,
&event.data.request_body_ref,
event.data.request_body_state,
),
(
&event.data.provider_request_body,
&event.data.provider_request_body_ref,
event.data.provider_request_body_state,
),
(
&event.data.response_body,
&event.data.response_body_ref,
event.data.response_body_state,
),
(
&event.data.client_response_body,
&event.data.client_response_body_ref,
event.data.client_response_body_state,
),
] {
assert!(body.is_none());
assert_eq!(actual_state, Some(state));
assert_eq!(reference.as_deref(), Some("usage://stale/reference"));
}
assert!(event
.data
.request_metadata
.as_ref()
.and_then(|value| value.get("body_capture"))
.is_none());
}
#[test]
fn from_stream_fields_and_clone_budget_preserve_typed_clear_with_residual_bodies() {
for state in [
UsageBodyCaptureState::None,
UsageBodyCaptureState::Disabled,
UsageBodyCaptureState::Unavailable,
] {
let mut source = captured_event();
source.data.request_metadata = None;
source.data.request_body_state = Some(state);
source.data.provider_request_body_state = Some(state);
source.data.response_body_state = Some(state);
source.data.client_response_body_state = Some(state);
source.data.request_body_ref = Some("usage://stale/reference".to_string());
source.data.provider_request_body_ref = Some("usage://stale/reference".to_string());
source.data.response_body_ref = Some("usage://stale/reference".to_string());
source.data.client_response_body_ref = Some("usage://stale/reference".to_string());
let fields = source.to_stream_fields().expect("wire serialization");
let decoded_budget = Arc::new(EventCaptureMemoryBudget::new(0));
let decoded = UsageEvent::from_stream_fields_with_capture_budget(
&fields,
Arc::clone(&decoded_budget),
)
.expect("wire decode");
assert_typed_body_clear_is_preserved(&decoded, state);
assert_eq!(decoded_budget.retained_bytes(), 0);
let clone_budget = Arc::new(EventCaptureMemoryBudget::new(
source.data.capture_heap_estimate(),
));
crate::runtime::prepare_event_capture_memory(&mut source, Arc::clone(&clone_budget));
let cloned = source.clone();
assert_typed_body_clear_is_preserved(&cloned, state);
assert!(source.data.request_body.is_some());
assert_eq!(clone_budget.downgraded_total(), 1);
drop(source);
assert_eq!(clone_budget.retained_bytes(), 0);
}
}
#[test]
fn from_stream_fields_legacy_metadata_only_facts_survive_before_billing() {
for limit in [0, 8192] {
for include_response in [false, true] {
let mut source = legacy_metadata_only_event();
if include_response {
source.data.response_body = Some(json!({"result": "response capture"}));
}
let fields = source
.to_stream_fields()
.expect("legacy wire serialization");
let budget = Arc::new(EventCaptureMemoryBudget::new(limit));
let decoded = UsageEvent::from_stream_fields_with_capture_budget(
&fields,
Arc::clone(&budget),
)
.expect("legacy wire decode");
// The worker enriches this clone before DTO conversion. Missing legacy bodies
// must not erase a previously derived TTL or turn an explicit zero into unknown.
let billing_event = decoded.clone();
let metadata = billing_event
.data
.request_metadata
.as_ref()
.expect("legacy facts");
assert_eq!(metadata["requested_reasoning_effort"], "high");
assert_eq!(metadata["provider_reasoning_effort"], "medium");
assert_eq!(metadata["provider_service_tier"], "priority");
assert_eq!(metadata["provider_actual_service_tier"], "default");
assert_eq!(metadata["provider_cache_ttl_minutes"], 60);
assert_eq!(billing_event.data.input_tokens, Some(0));
assert_eq!(billing_event.data.output_tokens, Some(0));
assert_eq!(billing_event.data.total_tokens, Some(0));
assert_eq!(billing_event.data.cache_read_input_tokens, Some(0));
assert_eq!(billing_event.data.cache_creation_input_tokens, Some(0));
assert_eq!(billing_event.data.actual_total_cost_usd, Some(0.0));
assert_eq!(billing_event.data.request_body_state, None);
assert_eq!(billing_event.data.provider_request_body_state, None);
drop((billing_event, decoded));
assert_eq!(budget.retained_bytes(), 0);
}
}
}
#[test]
fn from_stream_fields_typed_none_still_clears_metadata_only_request_facts() {
for limit in [0, 8192] {
let mut source = legacy_metadata_only_event();
source.data.request_body_state = Some(UsageBodyCaptureState::None);
source.data.provider_request_body_state = Some(UsageBodyCaptureState::None);
let fields = source.to_stream_fields().expect("wire serialization");
let budget = Arc::new(EventCaptureMemoryBudget::new(limit));
let decoded =
UsageEvent::from_stream_fields_with_capture_budget(&fields, Arc::clone(&budget))
.expect("wire decode");
let metadata = decoded
.data
.request_metadata
.as_ref()
.expect("response facts remain");
for key in [
"requested_reasoning_effort",
"provider_reasoning_effort",
"provider_service_tier",
"provider_cache_ttl_minutes",
] {
assert!(metadata.get(key).is_none(), "typed none must clear {key}");
}
assert_eq!(metadata["provider_actual_service_tier"], "default");
assert_eq!(
decoded.data.request_body_state,
Some(UsageBodyCaptureState::None)
);
assert_eq!(
decoded.data.provider_request_body_state,
Some(UsageBodyCaptureState::None)
);
assert_eq!(decoded.data.cache_read_input_tokens, Some(0));
assert_eq!(budget.retained_bytes(), 0);
assert_eq!(budget.downgraded_total(), 0);
}
}
fn legacy_metadata_only_event() -> UsageEvent {
UsageEvent::new(
UsageEventType::Completed,
"legacy-metadata-only",
UsageEventData {
provider_name: "openai".to_string(),
model: "gpt-5.6-sol".to_string(),
endpoint_api_format: Some("openai:responses".to_string()),
input_tokens: Some(0),
output_tokens: Some(0),
total_tokens: Some(0),
cache_read_input_tokens: Some(0),
cache_creation_input_tokens: Some(0),
actual_total_cost_usd: Some(0.0),
request_metadata: Some(json!({
"requested_reasoning_effort": "high",
"provider_reasoning_effort": "medium",
"provider_service_tier": "priority",
"provider_actual_service_tier": "default",
"provider_cache_ttl_minutes": 60
})),
..UsageEventData::default()
},
)
}
#[test]
fn from_stream_fields_body_omission_keeps_unknown_usage_unknown() {
let budget = Arc::new(EventCaptureMemoryBudget::new(0));
for event_type in [
UsageEventType::Completed,
UsageEventType::Failed,
UsageEventType::Cancelled,
] {
let event = UsageEvent::new(
event_type,
"usage-unavailable",
UsageEventData {
provider_name: "openai".to_string(),
model: "gpt-5".to_string(),
response_body: Some(json!({"error": "usage unavailable"})),
request_metadata: Some(json!({
"usage_available": false,
"usage_pricing_available": false
})),
..UsageEventData::default()
},
);
let fields = event.to_stream_fields().expect("wire serialization");
let decoded =
UsageEvent::from_stream_fields_with_capture_budget(&fields, Arc::clone(&budget))
.expect("wire decode");
assert_eq!(decoded.event_type, event_type);
assert_eq!(decoded.data.input_tokens, None);
assert_eq!(decoded.data.output_tokens, None);
assert_eq!(decoded.data.total_tokens, None);
assert_eq!(decoded.data.cache_read_input_tokens, None);
assert_eq!(decoded.data.cache_creation_input_tokens, None);
assert_eq!(decoded.data.actual_total_cost_usd, None);
assert_eq!(
decoded.data.request_metadata.as_ref().expect("metadata")["usage_available"],
false
);
assert_eq!(
decoded.data.request_metadata.as_ref().expect("metadata")
["usage_pricing_available"],
false
);
assert_eq!(
decoded.data.response_body_state,
Some(UsageBodyCaptureState::Truncated)
);
}
assert_eq!(budget.retained_bytes(), 0);
assert_eq!(budget.downgraded_total(), 3);
}
#[test]
fn from_stream_fields_invalid_envelopes_do_not_reserve_capture_memory() {
let budget = Arc::new(EventCaptureMemoryBudget::new(1024));
let mut unsupported = captured_event()
.to_stream_fields()
.expect("wire serialization");
let mut payload: serde_json::Value =
serde_json::from_str(&unsupported["payload"]).expect("json");
payload["v"] = json!(99);
unsupported.insert("payload".to_string(), payload.to_string());
for fields in [
BTreeMap::new(),
BTreeMap::from([("payload".to_string(), "not json".to_string())]),
unsupported,
] {
assert!(matches!(
UsageEvent::from_stream_fields_with_capture_budget(&fields, Arc::clone(&budget)),
Err(DataLayerError::UnexpectedValue(_))
));
assert_eq!(budget.retained_bytes(), 0);
assert_eq!(budget.downgraded_total(), 0);
}
}
#[test]
fn event_capture_budget_basic_policy_needs_no_diagnostic_allocation() {
let mut event = captured_event();
let budget = Arc::new(EventCaptureMemoryBudget::new(0));
apply_usage_body_capture_policy_to_event(UsageBodyCapturePolicy::default(), &mut event);
event.data.apply_capture_memory_budget(Arc::clone(&budget));
assert_eq!(
event.data.response_body_state,
Some(UsageBodyCaptureState::Disabled)
);
assert_eq!(event.data.total_tokens, Some(600));
assert_eq!(budget.retained_bytes(), 0);
assert_eq!(budget.downgraded_total(), 0);
}
#[test]
fn usage_event_round_trips_through_stream_fields() {
let event = UsageEvent::new(
@@ -0,0 +1,102 @@
use std::sync::{Arc, LazyLock};
#[cfg(test)]
use serde_json::Value;
#[doc(hidden)]
pub use aether_data_contracts::repository::usage::UsageCaptureRetention as UsageEventCaptureRetention;
pub(crate) use aether_data_contracts::repository::usage::{
usage_json_heap_estimate as json_heap_estimate,
UsageCaptureMemoryBudget as EventCaptureMemoryBudget,
};
const DEFAULT_CAPTURE_MEMORY_BUDGET_BYTES: usize = 128 * 1024 * 1024;
static CAPTURE_MEMORY_BUDGET: LazyLock<Arc<EventCaptureMemoryBudget>> = LazyLock::new(|| {
let limit = std::env::var("AETHER_USAGE_EVENT_CAPTURE_MEMORY_BUDGET_BYTES")
.ok()
.and_then(|value| value.trim().parse().ok())
.unwrap_or(DEFAULT_CAPTURE_MEMORY_BUDGET_BYTES);
Arc::new(EventCaptureMemoryBudget::new(limit))
});
pub(crate) fn shared_capture_memory_budget() -> Arc<EventCaptureMemoryBudget> {
Arc::clone(&CAPTURE_MEMORY_BUDGET)
}
pub(crate) fn capture_memory_metrics() -> (usize, usize, u64) {
CAPTURE_MEMORY_BUDGET.snapshot()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn event_capture_budget_resize_and_drop_release_estimate() {
let budget = Arc::new(EventCaptureMemoryBudget::new(16));
let mut retention = UsageEventCaptureRetention::default();
assert!(retention.reserve(Arc::clone(&budget), 12));
assert!(!retention.reserve(Arc::clone(&budget), 17));
assert_eq!(budget.retained_bytes(), 12);
assert_eq!(budget.downgraded_total(), 1);
assert!(retention.reserve(Arc::clone(&budget), 4));
assert_eq!(budget.retained_bytes(), 4);
drop(retention);
assert_eq!(budget.retained_bytes(), 0);
}
#[test]
fn event_capture_budget_unmanaged_clone_skips_estimation_and_empty_clone_is_free() {
let unmanaged = UsageEventCaptureRetention::default();
let (_, retained) =
unmanaged.clone_for_bodies(|| panic!("unmanaged JSON must not be scanned"));
assert!(retained);
let budget = Arc::new(EventCaptureMemoryBudget::new(0));
let mut managed = UsageEventCaptureRetention::default();
assert!(managed.reserve(Arc::clone(&budget), 0));
let (cloned, retained) = managed.clone_for_bodies(|| 0);
assert!(retained);
drop((managed, cloned));
assert_eq!(budget.retained_bytes(), 0);
assert_eq!(budget.downgraded_total(), 0);
}
#[test]
fn event_capture_budget_estimate_counts_string_and_array_spare_capacity() {
let mut text = String::with_capacity(1024);
text.push('x');
let mut array = Vec::with_capacity(16);
let expected = text.capacity() + array.capacity() * std::mem::size_of::<Value>();
array.push(Value::String(text));
assert_eq!(json_heap_estimate(&Value::Array(array)), expected);
}
#[test]
fn event_capture_budget_parallel_owners_never_exceed_shared_limit() {
let budget = Arc::new(EventCaptureMemoryBudget::new(1024));
let barrier = Arc::new(std::sync::Barrier::new(8));
std::thread::scope(|scope| {
for _ in 0..8 {
let budget = Arc::clone(&budget);
let barrier = Arc::clone(&barrier);
scope.spawn(move || {
for _ in 0..100 {
let mut retained = UsageEventCaptureRetention::default();
barrier.wait();
let _ = retained.reserve(Arc::clone(&budget), 400);
barrier.wait();
assert_eq!(budget.retained_bytes(), 800);
barrier.wait();
drop(retained);
barrier.wait();
assert_eq!(budget.retained_bytes(), 0);
barrier.wait();
}
});
}
});
assert_eq!(budget.retained_bytes(), 0);
assert_eq!(budget.downgraded_total(), 600);
}
}
@@ -0,0 +1,999 @@
use std::collections::BTreeMap;
use std::io::{self, Write};
use aether_data_contracts::repository::usage::{
resolve_provider_cache_ttl_minutes, UsageBodyCaptureState,
PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, PROVIDER_REASONING_EFFORT_METADATA_KEY,
PROVIDER_SERVICE_TIER_METADATA_KEY, REQUESTED_REASONING_EFFORT_METADATA_KEY,
};
use aether_data_contracts::DataLayerError;
use serde::ser::{Impossible, SerializeMap, SerializeStruct};
use serde::{Serialize, Serializer};
use serde_json::Value;
use super::{BorrowedUsageEventEnvelope, UsageEvent, UsageEventData, USAGE_EVENT_VERSION};
use crate::body_capture::mark_usage_event_capture_truncated;
use crate::request_metadata::{
attach_client_request_body_metadata, attach_provider_request_body_metadata,
attach_provider_response_body_metadata, clear_client_request_body_metadata,
clear_provider_request_body_metadata, request_body_derived_facts_action,
RequestBodyDerivedFactsAction,
};
const DIAGNOSTIC_FIELDS: [&str; 8] = [
"request_body",
"provider_request_body",
"response_body",
"client_response_body",
"request_headers",
"provider_request_headers",
"response_headers",
"client_response_headers",
];
const BODY_STATE_FIELDS: [&str; 4] = [
"request_body_state",
"provider_request_body_state",
"response_body_state",
"client_response_body_state",
];
const BODY_METADATA_KEYS: [&str; 4] =
["request", "provider_request", "response", "client_response"];
#[derive(Debug)]
pub(crate) struct EncodedUsageEvent {
pub(crate) fields: BTreeMap<String, String>,
pub(crate) diagnostics_omitted: bool,
}
pub(super) fn encode(
event: &UsageEvent,
max_bytes: usize,
) -> Result<EncodedUsageEvent, DataLayerError> {
let mut writer = BoundedJsonWriter::new(max_bytes);
if writer.serialize(&envelope(event, &event.data))? {
return writer.into_event(false);
}
// Conservatively reject oversized original metadata before cloning it, even
// when later fact normalization could make that metadata smaller.
let core = ProjectedData {
data: &event.data,
overrides: None,
};
if !writer.serialize(&envelope(event, &core))? {
return Err(wire_limit_error(max_bytes));
}
let overrides = WireOverrides::new(&event.data)?;
let projected = ProjectedData {
data: &event.data,
overrides: Some(&overrides),
};
if !writer.serialize(&envelope(event, &projected))? {
return Err(wire_limit_error(max_bytes));
}
writer.into_event(true)
}
fn envelope<'a, T: Serialize + ?Sized>(
event: &'a UsageEvent,
data: &'a T,
) -> BorrowedUsageEventEnvelope<'a, T> {
BorrowedUsageEventEnvelope {
v: USAGE_EVENT_VERSION,
event_type: event.event_type,
request_id: &event.request_id,
timestamp_ms: event.timestamp_ms,
data,
}
}
fn wire_limit_error(max_bytes: usize) -> DataLayerError {
DataLayerError::InvalidInput(format!(
"usage event exceeds the {max_bytes}-byte wire limit after omitting diagnostic bodies and headers"
))
}
struct BoundedJsonWriter {
bytes: Vec<u8>,
max_bytes: usize,
exceeded: bool,
}
impl BoundedJsonWriter {
fn new(max_bytes: usize) -> Self {
Self {
bytes: Vec::new(),
max_bytes,
exceeded: false,
}
}
fn serialize<T: Serialize + ?Sized>(&mut self, value: &T) -> Result<bool, DataLayerError> {
self.bytes.clear();
self.exceeded = false;
match serde_json::to_writer(&mut *self, value) {
Ok(()) => Ok(true),
Err(_) if self.exceeded => Ok(false),
Err(error) => Err(DataLayerError::UnexpectedValue(format!(
"failed to serialize usage event payload: {error}"
))),
}
}
fn into_event(self, diagnostics_omitted: bool) -> Result<EncodedUsageEvent, DataLayerError> {
let payload = String::from_utf8(self.bytes).map_err(|error| {
DataLayerError::UnexpectedValue(format!("usage event JSON was not UTF-8: {error}"))
})?;
Ok(EncodedUsageEvent {
fields: BTreeMap::from([("payload".to_string(), payload)]),
diagnostics_omitted,
})
}
}
impl Write for BoundedJsonWriter {
fn write(&mut self, bytes: &[u8]) -> io::Result<usize> {
if bytes.len() > self.max_bytes.saturating_sub(self.bytes.len()) {
self.exceeded = true;
return Err(io::Error::other("usage event wire limit exceeded"));
}
let required = self.bytes.len() + bytes.len();
if required > self.bytes.capacity() {
let capacity = required
.max(self.bytes.capacity().saturating_mul(2))
.min(self.max_bytes);
self.bytes
.try_reserve_exact(capacity - self.bytes.len())
.map_err(|error| {
io::Error::other(format!("usage event wire allocation failed: {error}"))
})?;
}
self.bytes.extend_from_slice(bytes);
Ok(bytes.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
struct WireOverrides {
truncated: [bool; 4],
metadata: Option<Value>,
}
impl WireOverrides {
fn new(data: &UsageEventData) -> Result<Self, DataLayerError> {
// The full v1 consumer decodes a JSON null Option<Value> as no body.
let request_body = data.request_body.as_ref().filter(|body| !body.is_null());
let provider_request_body = data
.provider_request_body
.as_ref()
.filter(|body| !body.is_null());
let body_cache_ttl = resolve_provider_cache_ttl_minutes(
data.endpoint_api_format
.as_deref()
.or(data.api_format.as_deref()),
data.target_model.as_deref().or(Some(data.model.as_str())),
Some(data.model.as_str()),
provider_request_body,
);
if body_cache_ttl.is_some()
&& data.provider_request_body_state == Some(UsageBodyCaptureState::None)
{
return Err(DataLayerError::InvalidInput(
"usage event cannot omit a provider request body whose cache TTL would be cleared by its explicit none capture state"
.to_string(),
));
}
let mut metadata = data.request_metadata.clone();
match request_body_derived_facts_action(request_body, data.request_body_state) {
RequestBodyDerivedFactsAction::Refresh => {
if request_body.is_some_and(|body| !body.is_object()) {
if let Some(Value::Object(object)) = metadata.as_mut() {
object.remove(REQUESTED_REASONING_EFFORT_METADATA_KEY);
}
} else {
metadata = attach_client_request_body_metadata(metadata, request_body);
}
}
RequestBodyDerivedFactsAction::Clear
if request_body.is_some() || data.request_body_state.is_some() =>
{
metadata = clear_client_request_body_metadata(metadata);
}
RequestBodyDerivedFactsAction::Clear | RequestBodyDerivedFactsAction::Preserve => {}
}
match request_body_derived_facts_action(
provider_request_body,
data.provider_request_body_state,
) {
RequestBodyDerivedFactsAction::Refresh => {
if provider_request_body.is_some_and(|body| !body.is_object()) {
// An authoritative scalar/array has no tier or reasoning,
// but billing still falls back to the metadata's cache TTL.
if let Some(Value::Object(object)) = metadata.as_mut() {
object.remove(PROVIDER_REASONING_EFFORT_METADATA_KEY);
object.remove(PROVIDER_SERVICE_TIER_METADATA_KEY);
}
} else {
metadata = attach_provider_request_body_metadata(
metadata,
data.endpoint_api_format
.as_deref()
.or(data.api_format.as_deref()),
data.target_model.as_deref().or(Some(data.model.as_str())),
Some(data.model.as_str()),
provider_request_body,
);
}
}
RequestBodyDerivedFactsAction::Clear
if provider_request_body.is_some()
|| data.provider_request_body_state.is_some() =>
{
metadata = clear_provider_request_body_metadata(metadata);
}
RequestBodyDerivedFactsAction::Clear | RequestBodyDerivedFactsAction::Preserve => {}
}
metadata = attach_provider_response_body_metadata(metadata, data.response_body.as_ref());
// Billing reads raw-body TTL before metadata regardless of capture state.
// Preserve that precedence independently of reasoning and tier authority.
if let Some(cache_ttl) = body_cache_ttl {
let object = metadata
.get_or_insert_with(|| Value::Object(serde_json::Map::new()))
.as_object_mut()
.ok_or_else(|| {
DataLayerError::InvalidInput(
"usage event cannot preserve provider cache TTL in non-object metadata after omitting diagnostic bodies"
.to_string(),
)
})?;
object.insert(
PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY.to_string(),
Value::Number(cache_ttl.into()),
);
}
let truncated = [
(data.request_body.as_ref(), data.request_body_state),
(
data.provider_request_body.as_ref(),
data.provider_request_body_state,
),
(data.response_body.as_ref(), data.response_body_state),
(
data.client_response_body.as_ref(),
data.client_response_body_state,
),
]
.map(|(body, state)| {
body.is_some_and(|body| !body.is_null())
&& !matches!(
state,
Some(
UsageBodyCaptureState::None
| UsageBodyCaptureState::Disabled
| UsageBodyCaptureState::Unavailable
)
)
});
for (truncated, key) in truncated.into_iter().zip(BODY_METADATA_KEYS) {
if !truncated {
continue;
}
mark_usage_event_capture_truncated(&mut metadata, key);
if let Some(entry) = metadata
.as_mut()
.and_then(|value| value.get_mut("body_capture"))
.and_then(|value| value.get_mut(key))
.and_then(Value::as_object_mut)
{
entry.insert(
"reason".to_string(),
Value::String("wire_limit_exceeded".to_string()),
);
}
}
Ok(Self {
truncated,
metadata,
})
}
}
struct ProjectedData<'a> {
data: &'a UsageEventData,
overrides: Option<&'a WireOverrides>,
}
impl Serialize for ProjectedData<'_> {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
self.data.serialize(FieldProjectionSerializer {
map: serializer.serialize_map(None)?,
overrides: self.overrides,
})
}
}
// Reuse UsageEventData's derived field traversal, including future fields and
// skip_serializing_if rules. Only diagnostic fields and explicit overrides differ.
struct FieldProjectionSerializer<'a, M> {
map: M,
overrides: Option<&'a WireOverrides>,
}
impl<M: SerializeMap> SerializeStruct for FieldProjectionSerializer<'_, M> {
type Ok = M::Ok;
type Error = M::Error;
fn serialize_field<T: Serialize + ?Sized>(
&mut self,
key: &'static str,
value: &T,
) -> Result<(), Self::Error> {
if DIAGNOSTIC_FIELDS.contains(&key) {
return Ok(());
}
if let Some(overrides) = self.overrides {
if key == "request_metadata"
|| BODY_STATE_FIELDS
.iter()
.zip(overrides.truncated)
.any(|(state, truncated)| *state == key && truncated)
{
return Ok(());
}
}
self.map.serialize_entry(key, value)
}
fn end(mut self) -> Result<Self::Ok, Self::Error> {
if let Some(overrides) = self.overrides {
for (key, truncated) in BODY_STATE_FIELDS.into_iter().zip(overrides.truncated) {
if truncated {
self.map
.serialize_entry(key, &UsageBodyCaptureState::Truncated)?;
}
}
if let Some(metadata) = overrides.metadata.as_ref() {
self.map.serialize_entry("request_metadata", metadata)?;
}
}
self.map.end()
}
}
fn expected_struct<E: serde::ser::Error, T>() -> Result<T, E> {
Err(E::custom("usage event data must serialize as a struct"))
}
macro_rules! reject_scalar_serialization {
($($name:ident($value:ident: $ty:ty)),* $(,)?) => {
$(fn $name(self, $value: $ty) -> Result<Self::Ok, Self::Error> {
let _ = $value;
expected_struct()
})*
};
}
impl<M: SerializeMap> Serializer for FieldProjectionSerializer<'_, M> {
type Ok = M::Ok;
type Error = M::Error;
type SerializeSeq = Impossible<Self::Ok, Self::Error>;
type SerializeTuple = Impossible<Self::Ok, Self::Error>;
type SerializeTupleStruct = Impossible<Self::Ok, Self::Error>;
type SerializeTupleVariant = Impossible<Self::Ok, Self::Error>;
type SerializeMap = Impossible<Self::Ok, Self::Error>;
type SerializeStruct = Self;
type SerializeStructVariant = Impossible<Self::Ok, Self::Error>;
reject_scalar_serialization! {
serialize_bool(value: bool), serialize_i8(value: i8), serialize_i16(value: i16),
serialize_i32(value: i32), serialize_i64(value: i64), serialize_i128(value: i128),
serialize_u8(value: u8), serialize_u16(value: u16), serialize_u32(value: u32),
serialize_u64(value: u64), serialize_u128(value: u128), serialize_f32(value: f32),
serialize_f64(value: f64), serialize_char(value: char), serialize_str(value: &str),
serialize_bytes(value: &[u8]),
}
fn serialize_none(self) -> Result<Self::Ok, Self::Error> {
expected_struct()
}
fn serialize_some<T: Serialize + ?Sized>(self, _: &T) -> Result<Self::Ok, Self::Error> {
expected_struct()
}
fn serialize_unit(self) -> Result<Self::Ok, Self::Error> {
expected_struct()
}
fn serialize_unit_struct(self, _: &'static str) -> Result<Self::Ok, Self::Error> {
expected_struct()
}
fn serialize_unit_variant(
self,
_: &'static str,
_: u32,
_: &'static str,
) -> Result<Self::Ok, Self::Error> {
expected_struct()
}
fn serialize_newtype_struct<T: Serialize + ?Sized>(
self,
_: &'static str,
_: &T,
) -> Result<Self::Ok, Self::Error> {
expected_struct()
}
fn serialize_newtype_variant<T: Serialize + ?Sized>(
self,
_: &'static str,
_: u32,
_: &'static str,
_: &T,
) -> Result<Self::Ok, Self::Error> {
expected_struct()
}
fn serialize_seq(self, _: Option<usize>) -> Result<Self::SerializeSeq, Self::Error> {
expected_struct()
}
fn serialize_tuple(self, _: usize) -> Result<Self::SerializeTuple, Self::Error> {
expected_struct()
}
fn serialize_tuple_struct(
self,
_: &'static str,
_: usize,
) -> Result<Self::SerializeTupleStruct, Self::Error> {
expected_struct()
}
fn serialize_tuple_variant(
self,
_: &'static str,
_: u32,
_: &'static str,
_: usize,
) -> Result<Self::SerializeTupleVariant, Self::Error> {
expected_struct()
}
fn serialize_map(self, _: Option<usize>) -> Result<Self::SerializeMap, Self::Error> {
expected_struct()
}
fn serialize_struct(
self,
_: &'static str,
_: usize,
) -> Result<Self::SerializeStruct, Self::Error> {
Ok(self)
}
fn serialize_struct_variant(
self,
_: &'static str,
_: u32,
_: &'static str,
_: usize,
) -> Result<Self::SerializeStructVariant, Self::Error> {
expected_struct()
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use serde_json::json;
use super::*;
use crate::event::UsageEventType;
use crate::event_capture_budget::EventCaptureMemoryBudget;
fn event() -> UsageEvent {
UsageEvent {
event_type: UsageEventType::Completed,
request_id: "wire-request".to_string(),
timestamp_ms: 1_234_567,
data: UsageEventData {
provider_name: "provider".to_string(),
model: "gpt-5.6-sol".to_string(),
endpoint_api_format: Some("openai:responses".to_string()),
user_id: Some("user-id".to_string()),
api_key_id: Some("key-id".to_string()),
provider_id: Some("provider-id".to_string()),
provider_endpoint_id: Some("endpoint-id".to_string()),
provider_api_key_id: Some("provider-key-id".to_string()),
input_tokens: Some(100),
output_tokens: Some(500),
total_tokens: Some(600),
cache_read_input_tokens: Some(0),
cache_creation_input_tokens: Some(25),
cache_creation_ephemeral_5m_input_tokens: Some(0),
cache_creation_ephemeral_1h_input_tokens: Some(25),
cache_read_cost_usd: Some(0.0),
total_cost_usd: Some(1.25),
actual_total_cost_usd: Some(1.25),
status_code: Some(200),
is_stream: Some(false),
candidate_index: Some(0),
first_byte_time_ms: Some(0),
error_message: Some(String::new()),
request_metadata: Some(json!({
"plan_usage_reservation_token": "550e8400-e29b-41d4-a716-446655440000",
"dimensions": {"image_count": 2, "size": "1024x1024", "quality": "high"},
"usage_available": true,
"usage_pricing_available": true
})),
..UsageEventData::default()
},
}
}
fn wire_value(encoded: &EncodedUsageEvent) -> Value {
serde_json::from_str(&encoded.fields["payload"]).expect("complete JSON wire payload")
}
#[test]
fn event_wire_exact_limit_accounts_for_json_escaping_and_utf8() {
let mut event = event();
event.request_id = "escaped\0\n\r\t\"\\\u{03bb}\u{1f600}".to_string();
let original = event.to_stream_fields().expect("original wire payload");
let length = original["payload"].len();
for limit in [length, length + 1] {
let encoded = event.to_bounded_stream_fields(limit).expect("exact fit");
assert_eq!(encoded.fields, original);
assert!(!encoded.diagnostics_omitted);
}
for limit in [0, 1, length - 1] {
assert!(matches!(
event.to_bounded_stream_fields(limit),
Err(DataLayerError::InvalidInput(_))
));
}
}
#[test]
fn event_wire_writer_does_not_append_past_limit_and_reuses_failed_buffer() {
let mut writer = BoundedJsonWriter::new(5);
writer.write_all(b"12345").expect("exact fit");
assert!(writer.write_all(b"6").is_err());
assert_eq!(writer.bytes, b"12345");
assert!(writer.exceeded);
assert!(!writer.serialize(&"\u{0000}").expect("size rejection"));
assert!(writer.bytes.len() <= 5);
assert!(writer.serialize(&"\u{03bb}").expect("valid UTF-8 retry"));
assert_eq!(writer.bytes, "\"\u{03bb}\"".as_bytes());
assert!(!writer.exceeded);
}
#[test]
fn event_wire_projection_transparently_forwards_derived_fields() {
#[derive(Serialize)]
struct FutureFields<'a> {
request_body: &'a Value,
new_billing_field: &'a Value,
zero_count: u64,
enabled: bool,
#[serde(skip_serializing_if = "Option::is_none")]
missing_field: Option<&'a str>,
}
let diagnostic = json!({"large": "x".repeat(8_192)});
let billing = json!({"nested": [0, false, "unchanged"]});
let data = FutureFields {
request_body: &diagnostic,
new_billing_field: &billing,
zero_count: 0,
enabled: false,
missing_field: None,
};
let mut bytes = Vec::new();
let mut serializer = serde_json::Serializer::new(&mut bytes);
data.serialize(FieldProjectionSerializer {
map: (&mut serializer)
.serialize_map(None)
.expect("object serializer"),
overrides: None,
})
.expect("project derived fields");
assert_eq!(
serde_json::from_slice::<Value>(&bytes).expect("projected JSON"),
json!({"new_billing_field": billing, "zero_count": 0, "enabled": false})
);
}
#[test]
fn event_wire_omission_preserves_billing_facts_refs_and_source_ownership() {
let mut event = event();
let padding = "x".repeat(16_384);
event.data.request_body = Some(json!({"reasoning": {"effort": "high"}, "input": padding}));
event.data.provider_request_body = Some(
json!({"model": "gpt-5.6-sol", "reasoning": {"effort": "medium"}, "service_tier": "priority", "input": padding}),
);
event.data.response_body = Some(json!({"service_tier": "Default", "output": padding}));
event.data.client_response_body = Some(json!({"output": padding}));
event.data.request_headers = Some(json!({"x-request": padding}));
event.data.provider_request_headers = Some(json!({"x-provider": padding}));
event.data.response_headers = Some(json!({"x-response": padding}));
event.data.client_response_headers = Some(json!({"x-client": padding}));
event.data.request_body_ref = Some("usage://wire-request/request_body".to_string());
event.data.request_body_state = Some(UsageBodyCaptureState::Inline);
event.data.provider_request_body_state = Some(UsageBodyCaptureState::Inline);
event.data.response_body_state = Some(UsageBodyCaptureState::Inline);
event.data.request_metadata.as_mut().unwrap()["body_capture"] = json!({
"response": {"state": "inline", "source_bytes": 123_456}
});
let weight = event.data.capture_heap_estimate();
let budget = Arc::new(EventCaptureMemoryBudget::new(weight));
event.data.apply_capture_memory_budget(Arc::clone(&budget));
let before = serde_json::to_value(&event).expect("source snapshot");
let original_wire: Value =
serde_json::from_str(&event.to_stream_fields().unwrap()["payload"]).unwrap();
let encoded = event
.to_bounded_stream_fields(8_192)
.expect("diagnostic omission");
assert!(encoded.diagnostics_omitted);
assert!(encoded.fields["payload"].len() <= 8_192);
let value = wire_value(&encoded);
for field in DIAGNOSTIC_FIELDS {
assert!(
value["data"].get(field).is_none(),
"{field} must be omitted"
);
}
for (field, original) in original_wire["data"].as_object().unwrap() {
if !DIAGNOSTIC_FIELDS.contains(&field.as_str())
&& !BODY_STATE_FIELDS.contains(&field.as_str())
&& field != "request_metadata"
{
assert_eq!(&value["data"][field], original, "core field {field}");
}
}
for (field, key) in BODY_STATE_FIELDS.into_iter().zip(BODY_METADATA_KEYS) {
assert_eq!(value["data"][field], "truncated");
assert_eq!(
value["data"]["request_metadata"]["body_capture"][key]["reason"],
"wire_limit_exceeded"
);
assert_eq!(
value["data"]["request_metadata"]["body_capture"][key]["stored_bytes"],
0
);
}
let metadata = &value["data"]["request_metadata"];
assert_eq!(metadata["requested_reasoning_effort"], "high");
assert_eq!(metadata["provider_reasoning_effort"], "medium");
assert_eq!(metadata["provider_service_tier"], "priority");
assert_eq!(metadata["provider_actual_service_tier"], "default");
assert_eq!(metadata["provider_cache_ttl_minutes"], 30);
assert_eq!(
metadata["body_capture"]["response"]["source_bytes"],
123_456
);
assert_eq!(
metadata["dimensions"],
before["data"]["request_metadata"]["dimensions"]
);
assert_eq!(serde_json::to_value(&event).unwrap(), before);
assert_eq!(budget.retained_bytes(), weight);
assert_eq!(budget.downgraded_total(), 0);
let decoded = UsageEvent::from_stream_fields(&encoded.fields).expect("wire decode");
let record = crate::build_upsert_usage_record_from_event(&decoded).expect("record mapping");
assert_eq!(record.input_tokens, Some(100));
assert_eq!(record.output_tokens, Some(500));
assert_eq!(record.cache_read_input_tokens, Some(0));
assert_eq!(record.cache_creation_ephemeral_5m_input_tokens, Some(0));
assert_eq!(record.cache_creation_ephemeral_1h_input_tokens, Some(25));
assert_eq!(record.error_message.as_deref(), Some(""));
assert_eq!(record.request_body_ref, event.data.request_body_ref);
assert_eq!(
record.request_body_state,
Some(UsageBodyCaptureState::Truncated)
);
let metadata = record.request_metadata.expect("record billing metadata");
assert_eq!(metadata["provider_cache_ttl_minutes"], 30);
assert_eq!(metadata["dimensions"]["image_count"], 2);
drop(event);
assert_eq!(budget.retained_bytes(), 0);
}
#[test]
fn event_wire_preserves_explicit_capture_states_and_legacy_metadata() {
for state in [
None,
Some(UsageBodyCaptureState::None),
Some(UsageBodyCaptureState::Disabled),
Some(UsageBodyCaptureState::Unavailable),
Some(UsageBodyCaptureState::Reference),
] {
let mut event = event();
event.data.request_body_state = state;
event.data.provider_request_body_state = state;
event.data.response_body_state = state;
event.data.client_response_body_state = state;
event.data.request_body_ref = Some("usage://wire-request/request_body".to_string());
event.data.response_headers = Some(json!({"large-header": "x".repeat(16_384)}));
event.data.request_metadata = Some(json!({
"requested_reasoning_effort": "high",
"provider_reasoning_effort": "medium",
"provider_service_tier": "priority",
"provider_cache_ttl_minutes": 60,
"provider_actual_service_tier": "flex"
}));
let before = serde_json::to_value(&event).unwrap();
let encoded = event
.to_bounded_stream_fields(4_096)
.expect("headers omitted");
let value = wire_value(&encoded);
for field in BODY_STATE_FIELDS {
assert_eq!(value["data"].get(field), before["data"].get(field));
}
let metadata = &value["data"]["request_metadata"];
if state == Some(UsageBodyCaptureState::None) {
assert!(metadata.get("requested_reasoning_effort").is_none());
assert!(metadata.get("provider_cache_ttl_minutes").is_none());
} else {
assert_eq!(metadata["requested_reasoning_effort"], "high");
assert_eq!(metadata["provider_cache_ttl_minutes"], 60);
}
assert_eq!(metadata["provider_actual_service_tier"], "flex");
assert_eq!(
value["data"]["request_body_ref"],
before["data"]["request_body_ref"]
);
assert_eq!(serde_json::to_value(&event).unwrap(), before);
if matches!(
state,
Some(
UsageBodyCaptureState::None
| UsageBodyCaptureState::Disabled
| UsageBodyCaptureState::Unavailable
)
) {
event.data.endpoint_api_format = Some("claude:messages".to_string());
event.data.request_body = Some(json!({"reasoning": {"effort": "low"}}));
event.data.provider_request_body =
Some(json!({"reasoning": {"effort": "low"}, "service_tier": "default"}));
event.data.response_body = Some(json!({"service_tier": "priority"}));
event.data.client_response_body = Some(json!({"stale": true}));
let baseline = UsageEvent::from_stream_fields_with_capture_budget(
&event.to_stream_fields().expect("full v1 fields"),
Arc::new(EventCaptureMemoryBudget::new(usize::MAX)),
)
.expect("full v1 consumer baseline");
let with_stale_bodies = wire_value(
&event
.to_bounded_stream_fields(4_096)
.expect("explicit capture states override stale bodies"),
);
for field in BODY_STATE_FIELDS {
assert_eq!(
with_stale_bodies["data"].get(field),
value["data"].get(field)
);
}
assert_eq!(
with_stale_bodies["data"]["request_metadata"],
serde_json::to_value(&baseline.data.request_metadata).unwrap(),
"wire omission must preserve the full v1 consumer's billing facts"
);
}
}
}
#[test]
fn event_wire_preserves_raw_body_cache_ttl_across_capture_states() {
use aether_data_contracts::repository::usage::extract_provider_cache_ttl_minutes_from_metadata;
for state in [
None,
Some(UsageBodyCaptureState::Inline),
Some(UsageBodyCaptureState::Reference),
Some(UsageBodyCaptureState::Truncated),
Some(UsageBodyCaptureState::Disabled),
Some(UsageBodyCaptureState::Unavailable),
] {
let mut event = event();
event.data.provider_request_body_state = state;
event.data.provider_request_body = Some(json!({
"prompt_cache_options": {"ttl": "30m"},
"service_tier": "default",
"reasoning": {"effort": "low"}
}));
event.data.response_headers = Some(json!({"large": "x".repeat(16_384)}));
event.data.request_metadata = Some(json!({
"provider_cache_ttl_minutes": 60,
"provider_service_tier": "priority",
"provider_reasoning_effort": "high"
}));
let before = serde_json::to_value(&event).unwrap();
let baseline = UsageEvent::from_stream_fields_with_capture_budget(
&event.to_stream_fields().expect("full v1 fields"),
Arc::new(EventCaptureMemoryBudget::new(usize::MAX)),
)
.expect("full v1 consumer baseline");
let encoded = event.to_bounded_stream_fields(4_096).expect("body omitted");
assert!(encoded.diagnostics_omitted);
let decoded = UsageEvent::from_stream_fields_with_capture_budget(
&encoded.fields,
Arc::new(EventCaptureMemoryBudget::new(usize::MAX)),
)
.expect("projected consumer event");
assert!(decoded.data.provider_request_body.is_none());
assert_eq!(
extract_provider_cache_ttl_minutes_from_metadata(
decoded.data.request_metadata.as_ref()
),
Some(30),
"raw body TTL must win over stale metadata for {state:?}"
);
for field in ["provider_service_tier", "provider_reasoning_effort"] {
assert_eq!(
decoded.data.request_metadata.as_ref().unwrap().get(field),
baseline.data.request_metadata.as_ref().unwrap().get(field),
"TTL preservation must not change {field} authority for {state:?}"
);
}
assert_eq!(serde_json::to_value(&event).unwrap(), before);
}
}
#[test]
fn event_wire_non_object_bodies_match_full_consumer_facts() {
use aether_data_contracts::repository::usage::{
extract_provider_cache_ttl_minutes_from_metadata,
resolve_provider_service_tier_from_request_capture,
};
for body in [
json!(null),
json!("opaque"),
json!([]),
json!(42),
json!(true),
] {
for state in [
None,
Some(UsageBodyCaptureState::None),
Some(UsageBodyCaptureState::Inline),
Some(UsageBodyCaptureState::Reference),
Some(UsageBodyCaptureState::Truncated),
Some(UsageBodyCaptureState::Disabled),
Some(UsageBodyCaptureState::Unavailable),
] {
let mut event = event();
event.data.request_body = Some(body.clone());
event.data.provider_request_body = Some(body.clone());
event.data.request_body_state = state;
event.data.provider_request_body_state = state;
event.data.response_headers = Some(json!({"large": "x".repeat(16_384)}));
event.data.request_metadata = Some(json!({
"requested_reasoning_effort": "medium",
"provider_reasoning_effort": "high",
"provider_service_tier": "priority",
"provider_cache_ttl_minutes": 60,
"unrelated": {"preserved": true}
}));
let original = event.to_stream_fields().expect("full v1 fields");
let baseline = UsageEvent::from_stream_fields_with_capture_budget(
&original,
Arc::new(EventCaptureMemoryBudget::new(usize::MAX)),
)
.expect("full v1 consumer baseline");
let encoded = event
.to_bounded_stream_fields(4_096)
.expect("omit diagnostics");
assert!(encoded.diagnostics_omitted);
let decoded = UsageEvent::from_stream_fields_with_capture_budget(
&encoded.fields,
Arc::new(EventCaptureMemoryBudget::new(usize::MAX)),
)
.expect("projected consumer event");
let tier = |data: &UsageEventData| {
resolve_provider_service_tier_from_request_capture(
data.provider_request_body.as_ref(),
data.provider_request_body_state,
data.request_metadata.as_ref(),
)
};
assert_eq!(
tier(&decoded.data),
tier(&baseline.data),
"provider tier for {body:?}, {state:?}"
);
let metadata = decoded.data.request_metadata.as_ref().unwrap();
let baseline_metadata = baseline.data.request_metadata.as_ref().unwrap();
assert_eq!(
extract_provider_cache_ttl_minutes_from_metadata(Some(metadata)),
extract_provider_cache_ttl_minutes_from_metadata(Some(baseline_metadata)),
"metadata TTL fallback for {body:?}, {state:?}"
);
let authoritative = !body.is_null()
&& matches!(
state,
None | Some(
UsageBodyCaptureState::Inline | UsageBodyCaptureState::Reference
)
);
for field in [
REQUESTED_REASONING_EFFORT_METADATA_KEY,
PROVIDER_REASONING_EFFORT_METADATA_KEY,
] {
assert_eq!(
metadata.get(field),
if authoritative {
None
} else {
baseline_metadata.get(field)
},
"reasoning authority for {body:?}, {state:?}, {field}"
);
}
assert_eq!(metadata["unrelated"], baseline_metadata["unrelated"]);
if body.is_null() {
assert_eq!(
decoded.data.request_body_state,
baseline.data.request_body_state
);
assert_eq!(
decoded.data.provider_request_body_state,
baseline.data.provider_request_body_state
);
}
assert_eq!(event.to_stream_fields().unwrap(), original);
}
}
}
#[test]
fn event_wire_rejects_cache_ttl_loss_from_explicit_none_capture_state() {
let mut event = event();
event.data.provider_request_body_state = Some(UsageBodyCaptureState::None);
event.data.provider_request_body = Some(json!({
"prompt_cache_options": {"ttl": "30m"},
"large": "x".repeat(16_384)
}));
let original = event.to_stream_fields().expect("full v1 fields");
let full = event
.to_bounded_stream_fields(original["payload"].len())
.expect("complete diagnostics remain representable");
assert_eq!(full.fields, original);
assert!(!full.diagnostics_omitted);
assert!(matches!(
event.to_bounded_stream_fields(4_096),
Err(DataLayerError::InvalidInput(_))
));
assert_eq!(event.to_stream_fields().unwrap(), original);
}
#[test]
fn event_wire_rejects_oversized_core_and_post_projection_metadata() {
for field in ["request_metadata", "error_message"] {
let mut event = event();
if field == "request_metadata" {
event.data.request_metadata = Some(json!({"large": "x".repeat(32_768)}));
} else {
event.data.error_message = Some("x".repeat(32_768));
}
let error = event
.to_bounded_stream_fields(1_024)
.expect_err("oversized core must fail");
assert!(matches!(error, DataLayerError::InvalidInput(_)));
assert!(
error.to_string().len() < 200,
"errors must not include payload contents"
);
}
let mut event = event();
event.data.response_body = Some(json!({"large": "x".repeat(8_192)}));
let core = ProjectedData {
data: &event.data,
overrides: None,
};
let core_size = serde_json::to_vec(&envelope(&event, &core)).unwrap().len();
assert!(
matches!(
event.to_bounded_stream_fields(core_size),
Err(DataLayerError::InvalidInput(_))
),
"added capture metadata must also fit the exact wire limit"
);
}
}
+30 -6
View File
@@ -1,5 +1,13 @@
use std::future::Future;
use std::sync::OnceLock;
use std::sync::{Mutex, OnceLock};
use std::time::Duration;
struct UsageBackgroundRuntime {
owner: Mutex<Option<tokio::runtime::Runtime>>,
handle: tokio::runtime::Handle,
}
static RUNTIME: OnceLock<UsageBackgroundRuntime> = OnceLock::new();
const DEFAULT_USAGE_BACKGROUND_RUNTIME_THREADS: usize = 8;
const MAX_USAGE_BACKGROUND_RUNTIME_THREADS: usize = 64;
@@ -18,12 +26,10 @@ where
F: Future + Send + 'static,
F::Output: Send + 'static,
{
usage_background_runtime().handle().spawn(task)
usage_background_runtime().handle.spawn(task)
}
fn usage_background_runtime() -> &'static tokio::runtime::Runtime {
static RUNTIME: OnceLock<&'static tokio::runtime::Runtime> = OnceLock::new();
fn usage_background_runtime() -> &'static UsageBackgroundRuntime {
RUNTIME.get_or_init(|| {
let worker_threads = usage_background_runtime_threads();
let runtime = tokio::runtime::Builder::new_multi_thread()
@@ -36,10 +42,28 @@ fn usage_background_runtime() -> &'static tokio::runtime::Runtime {
.thread_stack_size(USAGE_BACKGROUND_RUNTIME_STACK_BYTES)
.build()
.expect("usage background runtime should build");
Box::leak(Box::new(runtime))
UsageBackgroundRuntime {
handle: runtime.handle().clone(),
owner: Mutex::new(Some(runtime)),
}
})
}
/// Call outside Tokio after every UsageRuntime has drained. Does not start an unused runtime.
pub fn shutdown_usage_background_runtime(timeout: Duration) {
let Some(runtime) = RUNTIME.get() else {
return;
};
let owner = runtime
.owner
.lock()
.unwrap_or_else(|p| p.into_inner())
.take();
if let Some(owner) = owner {
owner.shutdown_timeout(timeout);
}
}
fn usage_background_runtime_threads() -> usize {
parse_usage_background_runtime_threads(
std::env::var(GATEWAY_USAGE_BACKGROUND_RUNTIME_THREADS_ENV)
+6
View File
@@ -1,15 +1,19 @@
mod body_capture;
pub mod config;
mod dead_letter_encoding;
pub mod event;
mod event_capture_budget;
mod executor;
mod keyed_lock;
pub mod queue;
mod queue_read_budget;
pub mod record;
pub mod report;
pub mod report_context;
mod request_metadata;
pub mod runtime;
pub mod settlement;
mod shutdown;
pub mod standardized_usage;
pub mod usage_mapper;
pub mod worker;
@@ -21,6 +25,7 @@ pub use body_capture::{
};
pub use config::UsageRuntimeConfig;
pub use event::{now_ms, UsageEvent, UsageEventData, UsageEventType, USAGE_EVENT_VERSION};
pub use executor::shutdown_usage_background_runtime;
pub use queue::UsageQueue;
pub use record::build_upsert_usage_record_from_event;
pub use report::{
@@ -50,6 +55,7 @@ pub use runtime::{
pub use settlement::{
reconcile_usage_policy_cost_for_event, settle_usage_if_needed, UsageSettlementWriter,
};
pub use shutdown::UsageProducerGuard;
pub use standardized_usage::StandardizedUsage;
pub use usage_mapper::{map_usage, map_usage_from_response, UsageMapper};
pub use worker::{
+556 -21
View File
@@ -1,14 +1,30 @@
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use serde_json::json;
use aether_data_contracts::DataLayerError;
use aether_runtime_state::{
RuntimeQueueEntry, RuntimeQueueReclaimConfig, RuntimeQueueStats, RuntimeQueueStore,
RuntimeQueueEntry, RuntimeQueueReclaimConfig, RuntimeQueueReclaimPage, RuntimeQueueStats,
RuntimeQueueStore, RuntimeQueueTransferOutcome,
};
use super::config::UsageRuntimeConfig;
use super::event::UsageEvent;
use super::event::{EncodedUsageEvent, UsageEvent};
use crate::dead_letter_encoding::{shared_dead_letter_encoding_budget, DeadLetterEncodingBudget};
use crate::queue_read_budget::{shared_queue_read_budget, QueueReadBudget, QueueReadReservation};
static PAYLOAD_DOWNGRADED_TOTAL: AtomicU64 = AtomicU64::new(0);
static PAYLOAD_REJECTED_TOTAL: AtomicU64 = AtomicU64::new(0);
pub(crate) fn payload_encoding_totals() -> (u64, u64) {
(
PAYLOAD_DOWNGRADED_TOTAL.load(Ordering::Relaxed),
PAYLOAD_REJECTED_TOTAL.load(Ordering::Relaxed),
)
}
pub(crate) fn is_permanent_enqueue_error(error: &DataLayerError) -> bool {
matches!(error, DataLayerError::InvalidInput(_))
}
#[derive(Clone)]
pub struct UsageQueue {
@@ -17,6 +33,35 @@ pub struct UsageQueue {
stream: String,
group: String,
dlq_stream: String,
read_budget: Arc<QueueReadBudget>,
dead_letter_encoding_budget: Arc<DeadLetterEncodingBudget>,
}
#[derive(Debug)]
pub(crate) enum UsageDeadLetterOutcome {
Transferred {
destination_id: String,
acked: usize,
},
Appended {
destination_id: String,
},
NotPending,
EncodingDeferred {
error: DataLayerError,
},
}
pub(crate) struct ReservedUsageReadBatch {
pub(crate) entries: Vec<RuntimeQueueEntry>,
pub(crate) requested_count: usize,
// Keep last: batch data must be dropped before its reservation is returned.
pub(crate) reservation: QueueReadReservation,
}
pub(crate) struct ReservedUsageReclaimPage {
pub(crate) page: RuntimeQueueReclaimPage,
pub(crate) reservation: QueueReadReservation,
}
impl UsageQueue {
@@ -31,9 +76,26 @@ impl UsageQueue {
group: config.consumer_group.clone(),
dlq_stream: config.dlq_stream_key.clone(),
config,
read_budget: shared_queue_read_budget(),
dead_letter_encoding_budget: shared_dead_letter_encoding_budget(),
})
}
#[cfg(test)]
pub(crate) fn with_read_budget(mut self, read_budget: Arc<QueueReadBudget>) -> Self {
self.read_budget = read_budget;
self
}
#[cfg(test)]
pub(crate) fn with_dead_letter_encoding_budget(
mut self,
budget: Arc<DeadLetterEncodingBudget>,
) -> Self {
self.dead_letter_encoding_budget = budget;
self
}
pub async fn ensure_consumer_group(&self) -> Result<(), DataLayerError> {
self.runner
.ensure_consumer_group(&self.stream, &self.group, "0-0")
@@ -41,12 +103,36 @@ impl UsageQueue {
}
pub async fn enqueue(&self, event: &UsageEvent) -> Result<String, DataLayerError> {
let fields = event.to_stream_fields()?;
let encoded = self.encode_event(event)?;
self.runner
.append_fields_with_maxlen(&self.stream, &fields, Some(self.config.stream_maxlen))
.append_fields_with_maxlen(
&self.stream,
&encoded.fields,
Some(self.config.stream_maxlen),
)
.await
}
pub(crate) fn validate_event(&self, event: &UsageEvent) -> Result<(), DataLayerError> {
self.encode_event(event).map(|_| ())
}
fn encode_event(&self, event: &UsageEvent) -> Result<EncodedUsageEvent, DataLayerError> {
let encoded = match event.to_bounded_stream_fields(self.config.queue_payload_max_bytes) {
Ok(encoded) => encoded,
Err(error) => {
if is_permanent_enqueue_error(&error) {
PAYLOAD_REJECTED_TOTAL.fetch_add(1, Ordering::Relaxed);
}
return Err(error);
}
};
if encoded.diagnostics_omitted {
PAYLOAD_DOWNGRADED_TOTAL.fetch_add(1, Ordering::Relaxed);
}
Ok(encoded)
}
pub async fn read_group(
&self,
consumer: &str,
@@ -62,13 +148,52 @@ impl UsageQueue {
.await
}
/// Workers retain this lease through processing. The public Vec API remains
/// compatible, but cannot preserve a reservation after returning its entries.
pub(crate) async fn read_group_reserved(
&self,
consumer: &str,
) -> Result<ReservedUsageReadBatch, DataLayerError> {
let (requested_count, mut reservation) = self
.read_budget
.reserve(
self.config.consumer_batch_size,
self.config.queue_payload_max_bytes,
)
.await?;
let entries = self
.runner
.read_group(
&self.stream,
&self.group,
consumer,
requested_count,
Some(self.config.consumer_block_ms.max(1)),
)
.await?;
reservation.observe_entries(&entries, self.config.queue_payload_max_bytes);
Ok(ReservedUsageReadBatch {
entries,
requested_count,
reservation,
})
}
pub async fn claim_stale(
&self,
consumer: &str,
start_id: &str,
) -> Result<Vec<RuntimeQueueEntry>, DataLayerError> {
Ok(self.claim_stale_page(consumer, start_id).await?.entries)
}
pub async fn claim_stale_page(
&self,
consumer: &str,
start_id: &str,
) -> Result<RuntimeQueueReclaimPage, DataLayerError> {
self.runner
.claim_stale(
.claim_stale_page(
&self.stream,
&self.group,
consumer,
@@ -81,10 +206,46 @@ impl UsageQueue {
.await
}
pub(crate) async fn claim_stale_page_reserved(
&self,
consumer: &str,
start_id: &str,
) -> Result<ReservedUsageReclaimPage, DataLayerError> {
let (requested_count, mut reservation) = self
.read_budget
.reserve(
self.config.reclaim_count,
self.config.queue_payload_max_bytes,
)
.await?;
let page = self
.runner
.claim_stale_page(
&self.stream,
&self.group,
consumer,
start_id,
RuntimeQueueReclaimConfig {
min_idle_ms: self.config.reclaim_idle_ms,
count: requested_count,
},
)
.await?;
reservation.observe_entries(&page.entries, self.config.queue_payload_max_bytes);
Ok(ReservedUsageReclaimPage { page, reservation })
}
pub async fn ack_and_delete(&self, ids: &[String]) -> Result<(), DataLayerError> {
self.runner.ack(&self.stream, &self.group, ids).await?;
self.ack_and_delete_counted(ids).await.map(|_| ())
}
pub(crate) async fn ack_and_delete_counted(
&self,
ids: &[String],
) -> Result<usize, DataLayerError> {
let acked = self.runner.ack(&self.stream, &self.group, ids).await?;
self.runner.delete(&self.stream, ids).await?;
Ok(())
Ok(acked)
}
pub async fn push_dead_letter(
@@ -92,20 +253,57 @@ impl UsageQueue {
entry: &RuntimeQueueEntry,
error: &str,
) -> Result<String, DataLayerError> {
let fields = std::collections::BTreeMap::from([(
"payload".to_string(),
serde_json::to_string(&json!({
"entry_id": entry.id,
"fields": entry.fields,
"error": error,
}))
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?,
)]);
let reservation = self.dead_letter_encoding_budget.try_reserve(entry, error)?;
let encoded = reservation
.encode_owned(entry.clone(), error.to_string())
.await?;
self.runner
.append_fields_with_maxlen(&self.dlq_stream, &fields, None)
.append_fields_with_maxlen(&self.dlq_stream, &encoded.fields, None)
.await
}
pub(crate) async fn transfer_dead_letter_owned(
&self,
entry: RuntimeQueueEntry,
error: String,
) -> Result<UsageDeadLetterOutcome, DataLayerError> {
let reservation = match self.dead_letter_encoding_budget.try_reserve(&entry, &error) {
Ok(reservation) => reservation,
Err(error) => return Ok(UsageDeadLetterOutcome::EncodingDeferred { error }),
};
let encoded = match reservation.encode_owned(entry, error).await {
Ok(encoded) => encoded,
Err(error) => return Ok(UsageDeadLetterOutcome::EncodingDeferred { error }),
};
match self
.runner
.try_transfer_pending_to_stream(
&self.stream,
&self.group,
&encoded.entry_id,
&self.dlq_stream,
&encoded.fields,
)
.await?
{
Some(RuntimeQueueTransferOutcome::Transferred {
destination_id,
acked,
..
}) => Ok(UsageDeadLetterOutcome::Transferred {
destination_id,
acked,
}),
Some(RuntimeQueueTransferOutcome::NotPending) => Ok(UsageDeadLetterOutcome::NotPending),
None => Ok(UsageDeadLetterOutcome::Appended {
destination_id: self
.runner
.append_fields_with_maxlen(&self.dlq_stream, &encoded.fields, None)
.await?,
}),
}
}
pub async fn stats(&self) -> Result<RuntimeQueueStats, DataLayerError> {
self.runner.stats(&self.stream, Some(&self.group)).await
}
@@ -136,10 +334,216 @@ fn usage_queue_runtime_settings(config: &UsageRuntimeConfig) -> UsageQueueRuntim
#[cfg(test)]
mod tests {
use super::{usage_queue_runtime_settings, UsageQueue, UsageQueueRuntimeSettings};
use super::{
usage_queue_runtime_settings, UsageDeadLetterOutcome, UsageQueue, UsageQueueRuntimeSettings,
};
use crate::dead_letter_encoding::DeadLetterEncodingBudget;
use crate::queue_read_budget::QueueReadBudget;
use crate::UsageRuntimeConfig;
use aether_runtime_state::{MemoryRuntimeStateConfig, RuntimeState};
use aether_data_contracts::DataLayerError;
use aether_runtime_state::{
MemoryRuntimeStateConfig, RuntimeQueueEntry, RuntimeQueueReclaimConfig, RuntimeQueueStats,
RuntimeQueueStore, RuntimeQueueTransferOutcome, RuntimeState,
};
use async_trait::async_trait;
use std::collections::BTreeMap;
use std::future::Future;
use std::sync::Arc;
use std::task::Poll;
use std::time::Duration;
struct HeldDeadLetterStore {
started: tokio::sync::Notify,
}
#[async_trait]
impl RuntimeQueueStore for HeldDeadLetterStore {
async fn ensure_consumer_group(
&self,
_stream: &str,
_group: &str,
_start_id: &str,
) -> Result<(), DataLayerError> {
unreachable!("store only exercises dead-letter writes")
}
async fn append_fields_with_maxlen(
&self,
_stream: &str,
fields: &BTreeMap<String, String>,
_maxlen: Option<usize>,
) -> Result<String, DataLayerError> {
assert!(fields.contains_key("payload"));
self.started.notify_one();
std::future::pending().await
}
async fn read_group(
&self,
_stream: &str,
_group: &str,
_consumer: &str,
_count: usize,
_block_ms: Option<u64>,
) -> Result<Vec<RuntimeQueueEntry>, DataLayerError> {
unreachable!("store only exercises dead-letter writes")
}
async fn claim_stale(
&self,
_stream: &str,
_group: &str,
_consumer: &str,
_start_id: &str,
_config: RuntimeQueueReclaimConfig,
) -> Result<Vec<RuntimeQueueEntry>, DataLayerError> {
unreachable!("store only exercises dead-letter writes")
}
async fn try_transfer_pending_to_stream(
&self,
_source: &str,
_group: &str,
_entry_id: &str,
_destination: &str,
fields: &BTreeMap<String, String>,
) -> Result<Option<RuntimeQueueTransferOutcome>, DataLayerError> {
assert!(fields.contains_key("payload"));
self.started.notify_one();
std::future::pending().await
}
async fn ack(
&self,
_stream: &str,
_group: &str,
_ids: &[String],
) -> Result<usize, DataLayerError> {
unreachable!("store only exercises dead-letter writes")
}
async fn delete(&self, _stream: &str, _ids: &[String]) -> Result<usize, DataLayerError> {
unreachable!("store only exercises dead-letter writes")
}
async fn stats(
&self,
_stream: &str,
_group: Option<&str>,
) -> Result<RuntimeQueueStats, DataLayerError> {
unreachable!("store only exercises dead-letter writes")
}
}
#[tokio::test]
async fn dead_letter_encoding_storage_wait_keeps_budget_until_public_or_owned_call_is_cancelled(
) {
for owned_transfer in [false, true] {
let store = Arc::new(HeldDeadLetterStore {
started: tokio::sync::Notify::new(),
});
let budget = Arc::new(DeadLetterEncodingBudget::new(4096, 1));
let queue = UsageQueue::new(store.clone(), UsageRuntimeConfig::default())
.unwrap()
.with_dead_letter_encoding_budget(Arc::clone(&budget));
let entry = RuntimeQueueEntry {
id: "1-0".to_string(),
fields: BTreeMap::from([("payload".to_string(), "original fields".to_string())]),
};
let task = tokio::spawn(async move {
if owned_transfer {
queue
.transfer_dead_letter_owned(entry, "failure".to_string())
.await
.map(|_| ())
} else {
queue.push_dead_letter(&entry, "failure").await.map(|_| ())
}
});
tokio::time::timeout(Duration::from_secs(2), store.started.notified())
.await
.unwrap();
assert_eq!(budget.snapshot().encoded_total, 1);
assert_eq!(budget.snapshot().active_jobs, 1);
assert!(budget.snapshot().reserved_bytes > 0);
task.abort();
assert!(matches!(task.await, Err(error) if error.is_cancelled()));
assert_eq!(budget.snapshot().active_jobs, 0);
assert_eq!(budget.snapshot().reserved_bytes, 0);
}
}
#[tokio::test]
async fn dead_letter_encoding_owned_defers_oversize_but_preserves_public_and_store_errors() {
let runner = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let queue = UsageQueue::new(runner, UsageRuntimeConfig::default())
.unwrap()
.with_dead_letter_encoding_budget(Arc::new(DeadLetterEncodingBudget::new(1, 1)));
let entry = RuntimeQueueEntry {
id: "1-0".to_string(),
fields: BTreeMap::from([("payload".to_string(), "original fields".to_string())]),
};
assert!(matches!(
queue
.transfer_dead_letter_owned(entry.clone(), "failure".to_string())
.await,
Ok(UsageDeadLetterOutcome::EncodingDeferred {
error: DataLayerError::InvalidInput(_),
})
));
assert!(matches!(
queue.push_dead_letter(&entry, "failure").await,
Err(DataLayerError::InvalidInput(_))
));
let budget = Arc::new(DeadLetterEncodingBudget::new(4096, 1));
let queue = queue.with_dead_letter_encoding_budget(Arc::clone(&budget));
// The absent source causes a native store error after successful encoding.
assert!(matches!(
queue
.transfer_dead_letter_owned(entry, "failure".to_string())
.await,
Err(DataLayerError::InvalidInput(_))
));
assert_eq!(budget.snapshot().encoded_total, 1);
assert_eq!(budget.snapshot().active_jobs, 0);
assert_eq!(budget.snapshot().reserved_bytes, 0);
assert_eq!(queue.dlq_stats().await.unwrap().stream_length, 0);
}
#[tokio::test]
async fn dead_letter_encoding_owned_defers_capacity_without_starting_store_work() {
let store = Arc::new(HeldDeadLetterStore {
started: tokio::sync::Notify::new(),
});
let budget = Arc::new(DeadLetterEncodingBudget::new(4096, 1));
let queue = UsageQueue::new(store, UsageRuntimeConfig::default())
.unwrap()
.with_dead_letter_encoding_budget(Arc::clone(&budget));
let entry = RuntimeQueueEntry {
id: "1-0".to_string(),
fields: BTreeMap::from([("payload".to_string(), "original fields".to_string())]),
};
let held = budget.try_reserve(&entry, "failure").unwrap();
let result = tokio::time::timeout(
Duration::from_secs(2),
queue.transfer_dead_letter_owned(entry, "failure".to_string()),
)
.await
.unwrap();
assert!(matches!(
result,
Ok(UsageDeadLetterOutcome::EncodingDeferred {
error: DataLayerError::TimedOut(_),
})
));
assert_eq!(budget.snapshot().encoded_total, 0);
assert_eq!(budget.snapshot().active_jobs, 1);
assert_eq!(budget.snapshot().capacity_rejected_total, 1);
drop(held);
assert_eq!(budget.snapshot().active_jobs, 0);
assert_eq!(budget.snapshot().reserved_bytes, 0);
}
#[test]
fn usage_queue_applies_runtime_block_and_batch_settings() {
@@ -164,4 +568,135 @@ mod tests {
}
);
}
fn reserved_test_queue(runner: Arc<RuntimeState>, budget: Arc<QueueReadBudget>) -> UsageQueue {
UsageQueue::new(
runner,
UsageRuntimeConfig {
enabled: true,
queue_payload_max_bytes: 8,
consumer_batch_size: 128,
consumer_block_ms: 1,
reclaim_count: 128,
reclaim_idle_ms: 1,
..UsageRuntimeConfig::default()
},
)
.expect("test queue")
.with_read_budget(budget)
}
#[tokio::test]
async fn queue_read_budget_read_and_reclaim_share_a_reservation_across_clones() {
let runtime = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let budget = Arc::new(QueueReadBudget::new(16, 16));
let queue = reserved_test_queue(Arc::clone(&runtime), Arc::clone(&budget));
let other = queue.clone();
queue.ensure_consumer_group().await.unwrap();
for _ in 0..6 {
runtime
.append_fields_with_maxlen(
&queue.stream,
&BTreeMap::from([("payload".to_string(), "12345678".to_string())]),
None,
)
.await
.unwrap();
}
let first = queue.read_group_reserved("reader").await.unwrap();
assert_eq!(first.requested_count, 2);
assert_eq!(first.entries.len(), 2);
let first_ids = first
.entries
.iter()
.map(|entry| entry.id.clone())
.collect::<Vec<_>>();
let extra_pending = runtime
.read_group(&queue.stream, &queue.group, "previous-reader", 2, None)
.await
.unwrap();
assert_eq!(extra_pending.len(), 2);
let next_cursor = extra_pending[0].id.clone();
drop(extra_pending);
assert_eq!(budget.snapshot().reserved_bytes, 16);
let mut reclaim = Box::pin(other.claim_stale_page_reserved("reclaimer", "0-0"));
std::future::poll_fn(|cx| {
assert!(reclaim.as_mut().poll(cx).is_pending());
Poll::Ready(())
})
.await;
assert_eq!(budget.snapshot().waiters, 1);
tokio::time::sleep(Duration::from_millis(5)).await;
drop(first);
let claimed = reclaim.await.unwrap();
assert_eq!(claimed.page.entries.len(), 2);
assert_eq!(claimed.page.next_start_id, next_cursor);
assert_eq!(
claimed
.page
.entries
.iter()
.map(|entry| entry.id.clone())
.collect::<Vec<_>>(),
first_ids
);
assert_eq!(budget.snapshot().reserved_bytes, 16);
assert_eq!(budget.snapshot().waiters, 0);
let mut next_read = Box::pin(queue.read_group_reserved("reader"));
std::future::poll_fn(|cx| {
assert!(next_read.as_mut().poll(cx).is_pending());
Poll::Ready(())
})
.await;
drop(claimed);
let next = next_read.await.unwrap();
assert_eq!(next.entries.len(), 2);
assert!(next
.entries
.iter()
.all(|entry| !first_ids.contains(&entry.id)));
drop(next);
assert_eq!(budget.snapshot().reserved_bytes, 0);
assert_eq!(budget.snapshot().wait_total, 2);
}
#[tokio::test]
async fn queue_read_budget_errors_and_empty_pages_release_reservations() {
let runtime = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let budget = Arc::new(QueueReadBudget::new(16, 16));
let queue = reserved_test_queue(runtime, Arc::clone(&budget));
assert!(queue.read_group_reserved("reader").await.is_err());
assert_eq!(budget.snapshot().reserved_bytes, 0);
assert!(queue
.claim_stale_page_reserved("reader", "0-0")
.await
.is_err());
assert_eq!(budget.snapshot().reserved_bytes, 0);
queue.ensure_consumer_group().await.unwrap();
let empty = queue.read_group_reserved("reader").await.unwrap();
assert!(empty.entries.is_empty());
assert_eq!(budget.snapshot().reserved_bytes, 0);
let page = queue
.claim_stale_page_reserved("reader", "0-0")
.await
.unwrap();
assert!(page.page.entries.is_empty());
assert_eq!(page.page.next_start_id, "0-0");
assert_eq!(budget.snapshot().reserved_bytes, 0);
drop((empty, page));
assert_eq!(budget.snapshot().reserved_bytes, 0);
}
#[test]
fn queue_read_budget_new_queues_and_clones_share_process_budget() {
let runtime = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let first = UsageQueue::new(runtime.clone(), UsageRuntimeConfig::default()).unwrap();
let second = UsageQueue::new(runtime, UsageRuntimeConfig::default()).unwrap();
let cloned = first.clone();
assert!(Arc::ptr_eq(&first.read_budget, &second.read_budget));
assert!(Arc::ptr_eq(&first.read_budget, &cloned.read_budget));
}
}
@@ -0,0 +1,373 @@
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, LazyLock};
use aether_data_contracts::DataLayerError;
use aether_runtime_state::RuntimeQueueEntry;
use tokio::sync::{OwnedSemaphorePermit, Semaphore, TryAcquireError};
const DEFAULT_READ_PAYLOAD_BUDGET_BYTES: usize = 128 * 1024 * 1024;
const DEFAULT_READ_BATCH_PAYLOAD_BYTES: usize = 8 * 1024 * 1024;
static READ_BUDGET: LazyLock<Arc<QueueReadBudget>> = LazyLock::new(|| {
let limit = std::env::var("AETHER_USAGE_QUEUE_READ_PAYLOAD_BUDGET_BYTES").ok();
let batch = std::env::var("AETHER_USAGE_QUEUE_READ_BATCH_PAYLOAD_BYTES").ok();
Arc::new(QueueReadBudget::new(
configured_bytes(limit.as_deref(), DEFAULT_READ_PAYLOAD_BUDGET_BYTES),
configured_bytes(batch.as_deref(), DEFAULT_READ_BATCH_PAYLOAD_BYTES),
))
});
pub(crate) fn shared_queue_read_budget() -> Arc<QueueReadBudget> {
Arc::clone(&READ_BUDGET)
}
pub(crate) fn queue_read_budget_metrics() -> QueueReadBudgetSnapshot {
READ_BUDGET.snapshot()
}
fn maximum_budget_bytes() -> usize {
Semaphore::MAX_PERMITS.min(u32::MAX as usize)
}
fn configured_bytes(raw: Option<&str>, fallback: usize) -> usize {
raw.and_then(|raw| raw.trim().parse::<u128>().ok())
.filter(|value| *value > 0)
.map(|value| value.min(maximum_budget_bytes() as u128) as usize)
.unwrap_or(fallback.min(maximum_budget_bytes()))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct QueueReadBudgetSnapshot {
pub(crate) limit_bytes: usize,
pub(crate) batch_limit_bytes: usize,
pub(crate) reserved_bytes: usize,
pub(crate) waiters: usize,
pub(crate) wait_total: u64,
/// Observed field names plus values, without cloning the returned strings.
pub(crate) actual_field_bytes_total: u64,
pub(crate) oversized_entries_total: u64,
pub(crate) oversized_batches_total: u64,
}
/// A process-wide reservation based on the current producer payload limit.
/// Historical or externally written messages can exceed the estimate. Field names,
/// allocation capacity, RESP decoding, and decoded JSON are not an RSS bound here.
pub(crate) struct QueueReadBudget {
limit_bytes: usize,
batch_limit_bytes: usize,
permits: Arc<Semaphore>,
reserved_bytes: AtomicUsize,
waiters: AtomicUsize,
wait_total: AtomicU64,
actual_field_bytes_total: AtomicU64,
oversized_entries_total: AtomicU64,
oversized_batches_total: AtomicU64,
}
impl QueueReadBudget {
pub(crate) fn new(limit_bytes: usize, batch_limit_bytes: usize) -> Self {
let limit_bytes = limit_bytes.clamp(1, maximum_budget_bytes());
let batch_limit_bytes = batch_limit_bytes.clamp(1, limit_bytes);
Self {
limit_bytes,
batch_limit_bytes,
permits: Arc::new(Semaphore::new(limit_bytes)),
reserved_bytes: AtomicUsize::new(0),
waiters: AtomicUsize::new(0),
wait_total: AtomicU64::new(0),
actual_field_bytes_total: AtomicU64::new(0),
oversized_entries_total: AtomicU64::new(0),
oversized_batches_total: AtomicU64::new(0),
}
}
pub(crate) fn snapshot(&self) -> QueueReadBudgetSnapshot {
QueueReadBudgetSnapshot {
limit_bytes: self.limit_bytes,
batch_limit_bytes: self.batch_limit_bytes,
reserved_bytes: self.reserved_bytes.load(Ordering::Relaxed),
waiters: self.waiters.load(Ordering::Relaxed),
wait_total: self.wait_total.load(Ordering::Relaxed),
actual_field_bytes_total: self.actual_field_bytes_total.load(Ordering::Relaxed),
oversized_entries_total: self.oversized_entries_total.load(Ordering::Relaxed),
oversized_batches_total: self.oversized_batches_total.load(Ordering::Relaxed),
}
}
pub(crate) async fn reserve(
self: &Arc<Self>,
requested_count: usize,
payload_limit: usize,
) -> Result<(usize, QueueReadReservation), DataLayerError> {
if payload_limit == 0 || payload_limit > self.limit_bytes {
return Err(DataLayerError::InvalidConfiguration(format!(
"usage queue payload limit {payload_limit} must be positive and not exceed the {}-byte read payload budget",
self.limit_bytes
)));
}
// A single valid payload may exceed the preferred batch target, but never
// the total budget. Clamp before multiplying or converting to u32 permits.
let count = requested_count
.max(1)
.min((self.batch_limit_bytes / payload_limit).max(1));
let reserved_bytes = count * payload_limit;
let permits = reserved_bytes as u32;
let permit = match Arc::clone(&self.permits).try_acquire_many_owned(permits) {
Ok(permit) => permit,
Err(TryAcquireError::NoPermits) => {
self.wait_total.fetch_add(1, Ordering::Relaxed);
self.waiters.fetch_add(1, Ordering::Relaxed);
let _waiting = WaitingReservation { budget: self };
Arc::clone(&self.permits)
.acquire_many_owned(permits)
.await
.map_err(|_| closed_budget_error())?
}
Err(TryAcquireError::Closed) => return Err(closed_budget_error()),
};
self.reserved_bytes
.fetch_add(reserved_bytes, Ordering::Relaxed);
Ok((
count,
QueueReadReservation {
budget: Arc::clone(self),
reserved_bytes,
permit: Some(permit),
},
))
}
}
fn closed_budget_error() -> DataLayerError {
DataLayerError::InvalidConfiguration("usage queue read payload budget is closed".to_string())
}
struct WaitingReservation<'a> {
budget: &'a QueueReadBudget,
}
impl Drop for WaitingReservation<'_> {
fn drop(&mut self) {
self.budget.waiters.fetch_sub(1, Ordering::Relaxed);
}
}
// Deliberately not Clone: every concurrently retained batch needs its own lease.
pub(crate) struct QueueReadReservation {
budget: Arc<QueueReadBudget>,
reserved_bytes: usize,
permit: Option<OwnedSemaphorePermit>,
}
impl QueueReadReservation {
pub(crate) fn observe_entries(&mut self, entries: &[RuntimeQueueEntry], payload_limit: usize) {
let mut value_bytes = 0usize;
let mut field_bytes = 0usize;
let mut oversized_entries = 0u64;
for entry in entries {
let mut entry_value_bytes = 0usize;
for (key, value) in &entry.fields {
entry_value_bytes = entry_value_bytes.saturating_add(value.len());
field_bytes = field_bytes
.saturating_add(key.len())
.saturating_add(value.len());
}
value_bytes = value_bytes.saturating_add(entry_value_bytes);
oversized_entries += u64::from(entry_value_bytes > payload_limit);
}
self.budget.actual_field_bytes_total.fetch_add(
u64::try_from(field_bytes).unwrap_or(u64::MAX),
Ordering::Relaxed,
);
self.budget
.oversized_entries_total
.fetch_add(oversized_entries, Ordering::Relaxed);
if value_bytes > self.reserved_bytes {
self.budget
.oversized_batches_total
.fetch_add(1, Ordering::Relaxed);
}
// Shrink unused payload estimates. Never wait for an upgrade after reading
// an oversized historical batch: other batches may hold all remaining bytes.
self.shrink_to(value_bytes.min(self.reserved_bytes));
}
fn shrink_to(&mut self, retained_bytes: usize) {
let released = self.reserved_bytes.saturating_sub(retained_bytes);
if released == 0 {
return;
}
let permit = self
.permit
.as_mut()
.expect("positive reservation must hold a permit")
.split(released)
.expect("released bytes must belong to this reservation");
self.reserved_bytes -= released;
self.budget
.reserved_bytes
.fetch_sub(released, Ordering::Relaxed);
drop(permit);
}
}
impl Drop for QueueReadReservation {
fn drop(&mut self) {
self.budget
.reserved_bytes
.fetch_sub(self.reserved_bytes, Ordering::Relaxed);
}
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use std::future::Future;
use std::task::Poll;
use super::*;
#[tokio::test]
async fn queue_read_budget_cancelled_wait_preserves_current_reservations() {
let budget = Arc::new(QueueReadBudget::new(32, 16));
let (_, first) = budget.reserve(8, 8).await.unwrap();
let (_, second) = budget.reserve(8, 8).await.unwrap();
let mut pending = Box::pin(budget.reserve(1, 8));
std::future::poll_fn(|cx| {
assert!(pending.as_mut().poll(cx).is_pending());
Poll::Ready(())
})
.await;
assert_eq!(budget.snapshot().reserved_bytes, 32);
assert_eq!(budget.snapshot().waiters, 1);
drop(pending);
assert_eq!(budget.snapshot().waiters, 0);
assert_eq!(budget.snapshot().wait_total, 1);
drop(first);
let (_, replacement) = budget.reserve(2, 8).await.unwrap();
assert_eq!(budget.snapshot().reserved_bytes, 32);
drop((second, replacement));
assert_eq!(budget.snapshot().reserved_bytes, 0);
assert_eq!(budget.permits.available_permits(), 32);
}
#[tokio::test]
async fn queue_read_budget_counts_payload_values_and_observes_legacy_excess_without_waiting() {
let budget = Arc::new(QueueReadBudget::new(16, 16));
let (count, mut reservation) = budget.reserve(2, 8).await.unwrap();
assert_eq!(count, 2);
let entries = [RuntimeQueueEntry {
id: "1-0".to_string(),
fields: BTreeMap::from([("payload".to_string(), "x".repeat(8))]),
}];
reservation.observe_entries(&entries, 8);
assert_eq!(budget.snapshot().reserved_bytes, 8);
assert_eq!(budget.snapshot().actual_field_bytes_total, 15);
assert_eq!(budget.snapshot().oversized_entries_total, 0);
let (_, mut legacy) = budget.reserve(1, 8).await.unwrap();
let legacy_entries = [RuntimeQueueEntry {
id: "2-0".to_string(),
fields: BTreeMap::from([
("payload".to_string(), "x".repeat(8)),
("extra".to_string(), "y".repeat(24)),
]),
}];
legacy.observe_entries(&legacy_entries, 8);
assert_eq!(budget.snapshot().reserved_bytes, 16);
assert_eq!(budget.snapshot().actual_field_bytes_total, 59);
assert_eq!(budget.snapshot().oversized_entries_total, 1);
assert_eq!(budget.snapshot().oversized_batches_total, 1);
assert_eq!(budget.snapshot().wait_total, 0);
drop((reservation, legacy));
assert_eq!(budget.snapshot().reserved_bytes, 0);
}
#[tokio::test]
async fn queue_read_budget_empty_response_releases_all_reserved_bytes() {
let budget = Arc::new(QueueReadBudget::new(16, 16));
let (_, mut reservation) = budget.reserve(2, 8).await.unwrap();
reservation.observe_entries(&[], 8);
assert_eq!(budget.snapshot().reserved_bytes, 0);
assert_eq!(budget.permits.available_permits(), 16);
drop(reservation);
assert_eq!(budget.snapshot().reserved_bytes, 0);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn queue_read_budget_concurrent_shrink_and_drop_preserve_shared_capacity() {
let budget = Arc::new(QueueReadBudget::new(64, 16));
let barrier = Arc::new(tokio::sync::Barrier::new(16));
let mut tasks = tokio::task::JoinSet::new();
for _ in 0..16 {
let budget = Arc::clone(&budget);
let barrier = Arc::clone(&barrier);
tasks.spawn(async move {
barrier.wait().await;
for _ in 0..16 {
let (count, mut reservation) = budget.reserve(128, 8).await.unwrap();
assert_eq!(count, 2);
assert!(budget.snapshot().reserved_bytes <= 64);
reservation.observe_entries(
&[RuntimeQueueEntry {
id: "1-0".to_string(),
fields: BTreeMap::from([(
"payload".to_string(),
"12345678".to_string(),
)]),
}],
8,
);
tokio::task::yield_now().await;
drop(reservation);
}
});
}
tokio::time::timeout(std::time::Duration::from_secs(5), async {
while let Some(result) = tasks.join_next().await {
result.expect("reservation task");
}
})
.await
.expect("all shared reservations complete");
assert_eq!(budget.snapshot().reserved_bytes, 0);
assert_eq!(budget.snapshot().waiters, 0);
assert_eq!(budget.snapshot().actual_field_bytes_total, 16 * 16 * 15);
assert_eq!(budget.permits.available_permits(), 64);
}
#[tokio::test]
async fn queue_read_budget_large_values_cannot_overflow_or_wait_for_impossible_permits() {
let budget = Arc::new(QueueReadBudget::new(usize::MAX, usize::MAX));
let maximum = maximum_budget_bytes();
assert_eq!(budget.snapshot().limit_bytes, maximum);
let (count, reservation) = budget.reserve(usize::MAX, 1).await.unwrap();
assert_eq!(count, maximum);
assert_eq!(budget.snapshot().reserved_bytes, maximum);
drop(reservation);
assert!(matches!(
budget.reserve(1, maximum + 1).await,
Err(DataLayerError::InvalidConfiguration(_))
));
assert!(matches!(
budget.reserve(1, 0).await,
Err(DataLayerError::InvalidConfiguration(_))
));
let small_batch = Arc::new(QueueReadBudget::new(32, 4));
let (count, reservation) = small_batch.reserve(usize::MAX, 16).await.unwrap();
assert_eq!(count, 1);
assert_eq!(small_batch.snapshot().reserved_bytes, 16);
drop(reservation);
}
#[test]
fn queue_read_budget_env_uses_positive_defaults_and_caps_extreme_values() {
for raw in [None, Some(""), Some("0"), Some("-1"), Some("bad")] {
assert_eq!(configured_bytes(raw, 128), 128);
}
assert_eq!(configured_bytes(Some(" 42 "), 128), 42);
assert_eq!(
configured_bytes(Some(&u128::MAX.to_string()), 128),
maximum_budget_bytes()
);
}
}
+44
View File
@@ -182,6 +182,7 @@ pub fn build_upsert_usage_record_from_event(
finalized_at_unix_secs,
created_at_unix_ms: Some(now_unix_secs),
updated_at_unix_secs: now_unix_secs,
capture_retention: data.capture_retention,
})
}
@@ -262,6 +263,49 @@ mod tests {
use super::build_upsert_usage_record_from_event;
#[test]
fn capture_retention_follows_event_bodies_into_record_and_its_clones() {
use aether_data_contracts::repository::usage::{
usage_json_heap_estimate, UsageCaptureMemoryBudget,
};
use std::sync::Arc;
let body = serde_json::Value::String("retained diagnostic".repeat(8));
let estimate =
4 * (std::mem::size_of::<serde_json::Value>() + usage_json_heap_estimate(&body));
let budget = Arc::new(UsageCaptureMemoryBudget::new(3 * estimate));
let mut event = UsageEvent::new(
UsageEventType::Completed,
"retained-record",
UsageEventData {
provider_name: "provider".to_owned(),
model: "model".to_owned(),
input_tokens: Some(5),
output_tokens: Some(7),
cache_read_input_tokens: Some(0),
request_body: Some(body.clone()),
provider_request_body: Some(body.clone()),
response_body: Some(body.clone()),
client_response_body: Some(body),
..UsageEventData::default()
},
);
event.data.apply_capture_memory_budget(Arc::clone(&budget));
assert_eq!(budget.retained_bytes(), estimate);
let record = build_upsert_usage_record_from_event(&event).unwrap();
assert_eq!(budget.retained_bytes(), 2 * estimate);
drop(event);
assert_eq!(budget.retained_bytes(), estimate);
let cloned = record.clone();
assert_eq!(budget.retained_bytes(), 2 * estimate);
assert_eq!(cloned.cache_read_input_tokens, Some(0));
assert_eq!(cloned.response_body, record.response_body);
drop(record);
assert_eq!(budget.retained_bytes(), estimate);
drop(cloned);
assert_eq!(budget.retained_bytes(), 0);
}
#[test]
fn builds_upsert_record_from_terminal_event() {
let record = build_upsert_usage_record_from_event(&UsageEvent {
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,432 @@
use super::*;
use super::super::{run_usage_enqueue_retry_worker, TerminalPersistenceOutcome};
const PAYLOAD_LIMIT: usize = 4 * 1024;
fn payload_config(name: &str) -> UsageRuntimeConfig {
UsageRuntimeConfig {
enabled: true,
queue_terminal_events: true,
queue_lifecycle_events: true,
stream_key: format!("usage:events:test:payload:{name}"),
consumer_group: format!("usage_consumers_payload_{name}"),
queue_payload_max_bytes: PAYLOAD_LIMIT,
consumer_block_ms: 1,
enqueue_retry_buffer_capacity: 8,
enqueue_retry_workers: 1,
enqueue_retry_initial_backoff_ms: 1,
enqueue_retry_max_backoff_ms: 2,
..UsageRuntimeConfig::default()
}
}
fn payload_event(request_id: &str, oversized: bool) -> UsageEvent {
UsageEvent {
event_type: UsageEventType::Completed,
request_id: request_id.to_string(),
timestamp_ms: 123_000,
data: UsageEventData {
user_id: Some("user-payload".to_string()),
api_key_id: Some("key-payload".to_string()),
provider_name: "openai".to_string(),
provider_id: Some("provider-payload".to_string()),
provider_api_key_id: Some("provider-key-payload".to_string()),
// Model identity is a core field, so omitting diagnostic bodies cannot make this fit.
model: if oversized {
"m".repeat(PAYLOAD_LIMIT * 2)
} else {
"gpt-5".to_string()
},
target_model: Some("gpt-5".to_string()),
api_format: Some("openai:responses".to_string()),
endpoint_api_format: Some("openai:responses".to_string()),
input_tokens: Some(100),
output_tokens: Some(25),
total_tokens: Some(125),
cache_creation_input_tokens: Some(30),
cache_creation_ephemeral_5m_input_tokens: Some(0),
cache_creation_ephemeral_1h_input_tokens: Some(30),
cache_read_input_tokens: Some(0),
total_cost_usd: Some(0.5),
actual_total_cost_usd: Some(0.25),
status_code: Some(200),
error_message: Some("preserve error presence and text".to_string()),
first_byte_time_ms: Some(12),
response_time_ms: Some(34),
request_headers: Some(json!({"x-request": "original"})),
provider_request_headers: Some(json!({"x-provider-request": "original"})),
response_headers: Some(json!({"x-provider-response": "original"})),
client_response_headers: Some(json!({"x-client-response": "original"})),
request_body: Some(json!({"reasoning": {"effort": "high"}, "input": "original"})),
request_body_state: Some(UsageBodyCaptureState::Inline),
provider_request_body: Some(json!({
"model": "gpt-5", "service_tier": "priority",
"prompt_cache_retention": "24h", "input": "original provider input"
})),
provider_request_body_state: Some(UsageBodyCaptureState::Inline),
response_body: Some(json!({"service_tier": "default", "output": "original output"})),
response_body_state: Some(UsageBodyCaptureState::Inline),
client_response_body: Some(json!({"output": "original client output"})),
client_response_body_state: Some(UsageBodyCaptureState::Inline),
request_metadata: Some(json!({
"api_key_is_standalone": true,
"plan_usage_reservation_token": "550e8400-e29b-41d4-a716-446655440000",
"plan_usage_reservation_deferred": true,
"usage_available": true,
"usage_pricing_available": true,
"provider_cache_ttl_minutes": 1440,
"provider_service_tier": "priority",
"provider_actual_service_tier": "default",
"dimensions": {
"image_count": 2, "image_size": "1024x1024", "image_quality": "high",
"image_output_format": "png", "reasoning_tokens": 0
}
})),
..UsageEventData::default()
},
}
}
async fn payload_queue(config: &UsageRuntimeConfig) -> (UsageQueue, Arc<FlakyAppendQueueStore>) {
let inner: Arc<dyn RuntimeQueueStore> =
Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let runner = Arc::new(FlakyAppendQueueStore::new(inner, 0));
let queue = UsageQueue::new(runner.clone(), config.clone()).expect("payload queue");
queue.ensure_consumer_group().await.expect("payload group");
(queue, runner)
}
#[tokio::test]
async fn terminal_oversize_uses_original_event_for_direct_fallback_without_opening_circuit() {
let config = payload_config("terminal_direct");
let (queue, runner) = payload_queue(&config).await;
let store = CloneQueueConfiguredUsageStore {
records: Arc::new(Mutex::new(Vec::new())),
queue: runner.clone(),
};
let runtime = UsageRuntime::new(config).expect("usage runtime");
let event = payload_event("payload-terminal-direct", true);
let expected = crate::build_upsert_usage_record_from_event(&event).expect("original record");
assert!(matches!(
queue.enqueue(&event).await,
Err(DataLayerError::InvalidInput(_))
));
assert_eq!(runner.append_attempts.load(Ordering::Acquire), 0);
assert_eq!(
runtime.enqueue_or_write_terminal(&store, event).await,
TerminalPersistenceOutcome::PersistedDirectly
);
{
let records = store.records.lock().expect("records lock");
assert_eq!(records.as_slice(), &[expected]);
assert_eq!(records[0].cache_read_input_tokens, Some(0));
assert_eq!(records[0].cache_creation_ephemeral_5m_input_tokens, Some(0));
assert_eq!(
records[0].request_body_state,
Some(UsageBodyCaptureState::Inline)
);
assert!(records[0].provider_request_body.is_some());
assert!(records[0].request_headers.is_some());
}
assert_eq!(
runtime
.terminal_enqueue_state
.circuit_open_until_unix_ms
.load(Ordering::Acquire),
0
);
let snapshot = runtime.metrics_snapshot();
assert_eq!(snapshot.terminal_direct_fallback_succeeded_total, 1);
assert_eq!(snapshot.enqueue_retry_scheduled_total, 0);
assert_eq!(snapshot.enqueue_retry_pending, 0);
assert_eq!(
runtime
.enqueue_or_write_terminal(&store, payload_event("payload-after-direct", false))
.await,
TerminalPersistenceOutcome::Queued
);
assert_eq!(runner.successful_appends.load(Ordering::Acquire), 1);
assert_eq!(store.records.lock().expect("records lock").len(), 1);
}
#[tokio::test]
async fn terminal_oversize_direct_failure_preserves_first_byte_and_does_not_retry_or_open_circuit()
{
let config = payload_config("terminal_failed");
let (_, runner) = payload_queue(&config).await;
let store = FailingWriteQueueConfiguredUsageStore {
queue: runner.clone(),
upsert_attempts: Arc::new(AtomicUsize::new(0)),
};
let runtime = UsageRuntime::new(config).expect("usage runtime");
let request_id = "payload-terminal-failed";
let generation = runtime
.lifecycle_coalescer
.mark_first_byte(request_id)
.await
.expect("first byte");
assert_eq!(
runtime
.enqueue_or_write_terminal(&store, payload_event(request_id, true))
.await,
TerminalPersistenceOutcome::Failed
);
assert!(
runtime
.lifecycle_coalescer
.first_byte_is_current(request_id, generation)
.await
);
assert_eq!(runner.append_attempts.load(Ordering::Acquire), 0);
assert_eq!(store.upsert_attempts.load(Ordering::Acquire), 1);
assert_eq!(
runtime
.terminal_enqueue_state
.circuit_open_until_unix_ms
.load(Ordering::Acquire),
0
);
let snapshot = runtime.metrics_snapshot();
assert_eq!(snapshot.terminal_direct_fallback_failed_total, 1);
assert_eq!(snapshot.terminal_enqueue_deferred_dropped_total, 1);
assert_eq!(snapshot.enqueue_retry_scheduled_total, 0);
assert_eq!(snapshot.enqueue_retry_pending, 0);
assert_eq!(
runtime
.enqueue_or_write_terminal(&store, payload_event("payload-after-failed", false))
.await,
TerminalPersistenceOutcome::Queued
);
assert_eq!(runner.successful_appends.load(Ordering::Acquire), 1);
assert_eq!(store.upsert_attempts.load(Ordering::Acquire), 1);
}
#[tokio::test]
async fn terminal_oversize_direct_failure_stays_failed_when_primary_enqueue_is_deferred() {
for (name, circuit_open) in [("circuit_open", true), ("in_flight_limit", false)] {
let mut config = payload_config(name);
config.terminal_enqueue_max_in_flight = 1;
let (_, runner) = payload_queue(&config).await;
let store = FailingWriteQueueConfiguredUsageStore {
queue: runner.clone(),
upsert_attempts: Arc::new(AtomicUsize::new(0)),
};
let runtime = UsageRuntime::new(config).expect("usage runtime");
let request_id = format!("payload-terminal-{name}");
let generation = runtime
.lifecycle_coalescer
.mark_first_byte(&request_id)
.await
.expect("first byte");
let original_deadline = if circuit_open {
let deadline = super::super::now_unix_ms().saturating_add(60_000);
runtime.terminal_enqueue_state.open_circuit(deadline);
deadline
} else {
0
};
let held_guard = if circuit_open {
None
} else {
Some(
runtime
.terminal_enqueue_state
.try_acquire_in_flight(1)
.expect("hold the only enqueue slot"),
)
};
assert_eq!(
timeout(
Duration::from_secs(2),
runtime.enqueue_or_write_terminal(&store, payload_event(&request_id, true)),
)
.await
.expect("bounded terminal fallback"),
TerminalPersistenceOutcome::Failed,
"{name} must not report an oversized event as buffered"
);
assert!(
runtime
.lifecycle_coalescer
.first_byte_is_current(&request_id, generation)
.await
);
assert_eq!(runner.append_attempts.load(Ordering::Acquire), 0);
assert_eq!(store.upsert_attempts.load(Ordering::Acquire), 1);
assert_eq!(
runtime
.terminal_enqueue_state
.circuit_open_until_unix_ms
.load(Ordering::Acquire),
original_deadline
);
let snapshot = runtime.metrics_snapshot();
assert_eq!(snapshot.terminal_direct_fallback_failed_total, 1);
assert_eq!(snapshot.terminal_enqueue_deferred_dropped_total, 1);
assert_eq!(snapshot.terminal_enqueue_deferred_retry_total, 0);
assert_eq!(snapshot.enqueue_retry_permanent_failure_total, 1);
assert_eq!(snapshot.enqueue_retry_scheduled_total, 0);
assert_eq!(snapshot.enqueue_retry_pending, 0);
assert_eq!(
runtime.terminal_enqueue_state.in_flight(),
u64::from(!circuit_open)
);
drop(held_guard);
assert_eq!(runtime.terminal_enqueue_state.in_flight(), 0);
}
}
#[tokio::test]
async fn retry_worker_discards_oversize_and_drains_next_event_on_the_same_shard() {
let config = payload_config("retry_drain");
let (queue, runner) = payload_queue(&config).await;
let (sender, receiver) = mpsc::channel(2);
let dispatcher = UsageEnqueueRetryDispatcher {
senders: vec![sender],
metrics: Arc::new(Default::default()),
};
let metrics_view = UsageEnqueueRetryDispatcher {
senders: Vec::new(),
metrics: Arc::clone(&dispatcher.metrics),
};
for (request_id, oversized) in [
("payload-retry-oversize", true),
("payload-retry-small", false),
] {
// Bypass admission to exercise the worker's defense for an already buffered event.
assert!(dispatcher
.schedule_item(
queue.clone(),
payload_event(request_id, oversized),
"terminal",
Some("prior transient failure"),
)
.is_some());
}
assert_eq!(dispatcher.pending(), 2);
drop(dispatcher);
// Run the real worker as a cancellable future, so a regression cannot leave a detached retry.
timeout(
Duration::from_secs(2),
run_usage_enqueue_retry_worker(0, config, receiver, Arc::clone(&metrics_view.metrics)),
)
.await
.expect("the permanent failure must not block the shard");
assert_eq!(metrics_view.permanent_failure_total(), 1);
assert_eq!(metrics_view.recovered_total(), 1);
assert_eq!(metrics_view.pending(), 0);
assert_eq!(runner.append_attempts.load(Ordering::Acquire), 1);
let entries = queue
.read_group("payload-retry-reader")
.await
.expect("queue read");
assert_eq!(entries.len(), 1);
assert_eq!(
UsageEvent::from_stream_fields(&entries[0].fields)
.expect("queued event")
.request_id,
"payload-retry-small"
);
}
#[tokio::test]
async fn retry_dispatcher_rejects_oversize_for_permanent_and_transient_causes_without_consuming_slots(
) {
let config = payload_config("retry_reject");
let (queue, runner) = payload_queue(&config).await;
let (sender, receiver) = mpsc::channel(1);
let dispatcher = UsageEnqueueRetryDispatcher {
senders: vec![sender],
metrics: Arc::new(Default::default()),
};
let oversized = payload_event("payload-rejected", true);
let error = queue.enqueue(&oversized).await.expect_err("oversize input");
assert!(matches!(error, DataLayerError::InvalidInput(_)));
assert!(!dispatcher.schedule(queue.clone(), oversized, "terminal", error));
for cause in [
DataLayerError::TimedOut("primary enqueue was deferred".to_string()),
DataLayerError::Redis("prior transient failure".to_string()),
] {
assert!(!dispatcher.schedule(
queue.clone(),
payload_event("payload-rejected-transient-cause", true),
"terminal",
cause,
));
}
assert_eq!(dispatcher.permanent_failure_total(), 3);
assert_eq!(dispatcher.pending(), 0);
assert_eq!(dispatcher.scheduled_total(), 0);
assert!(dispatcher.schedule(
queue,
payload_event("payload-after-reject", false),
"terminal",
DataLayerError::Redis("retryable failure".to_string()),
));
let metrics_view = UsageEnqueueRetryDispatcher {
senders: Vec::new(),
metrics: Arc::clone(&dispatcher.metrics),
};
drop(dispatcher);
timeout(
Duration::from_secs(2),
run_usage_enqueue_retry_worker(0, config, receiver, Arc::clone(&metrics_view.metrics)),
)
.await
.expect("retry worker drain");
assert_eq!(metrics_view.permanent_failure_total(), 3);
assert_eq!(metrics_view.recovered_total(), 1);
assert_eq!(metrics_view.pending(), 0);
assert_eq!(runner.append_attempts.load(Ordering::Acquire), 1);
}
#[tokio::test]
async fn lifecycle_oversize_does_not_open_circuit_or_block_the_next_lifecycle_event() {
let config = payload_config("lifecycle");
let (queue, runner) = payload_queue(&config).await;
let store = CloneQueueConfiguredUsageStore {
records: Arc::new(Mutex::new(Vec::new())),
queue: runner.clone(),
};
let runtime = UsageRuntime::new(config).expect("usage runtime");
let mut oversized = payload_event("payload-lifecycle-oversize", true);
oversized.event_type = UsageEventType::Streaming;
assert!(!runtime.enqueue_lifecycle_event(&store, oversized).await);
assert_eq!(
runtime
.lifecycle_enqueue_state
.circuit_open_until_unix_ms
.load(Ordering::Acquire),
0
);
assert_eq!(runner.append_attempts.load(Ordering::Acquire), 0);
let snapshot = runtime.metrics_snapshot();
assert_eq!(snapshot.enqueue_retry_scheduled_total, 0);
assert_eq!(snapshot.enqueue_retry_pending, 0);
assert!(store.records.lock().expect("records lock").is_empty());
let mut small = payload_event("payload-lifecycle-small", false);
small.event_type = UsageEventType::Streaming;
assert!(runtime.enqueue_lifecycle_event(&store, small).await);
assert_eq!(runner.successful_appends.load(Ordering::Acquire), 1);
let entries = queue
.read_group("payload-lifecycle-reader")
.await
.expect("queue read");
assert_eq!(entries.len(), 1);
let queued = UsageEvent::from_stream_fields(&entries[0].fields).expect("queued lifecycle");
assert_eq!(queued.request_id, "payload-lifecycle-small");
assert_eq!(queued.event_type, UsageEventType::Streaming);
assert_eq!(queued.data.first_byte_time_ms, Some(12));
}
@@ -0,0 +1,366 @@
use super::*;
fn config() -> UsageRuntimeConfig {
UsageRuntimeConfig {
enabled: true,
queue_terminal_events: true,
queue_lifecycle_events: true,
worker_count: 2,
consumer_block_ms: 60_000,
enqueue_retry_buffer_capacity: 256,
enqueue_retry_initial_backoff_ms: 60_000,
enqueue_retry_max_backoff_ms: 60_000,
..UsageRuntimeConfig::default()
}
}
fn store() -> CloneQueueConfiguredUsageStore {
CloneQueueConfiguredUsageStore {
records: Arc::new(Mutex::new(Vec::new())),
queue: Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())),
}
}
fn terminal(request_id: &str) -> UsageEvent {
UsageEvent::new(
UsageEventType::Completed,
request_id,
UsageEventData {
provider_name: "openai".to_string(),
model: "test".to_string(),
status_code: Some(200),
input_tokens: Some(3),
output_tokens: Some(7),
total_tokens: Some(10),
..UsageEventData::default()
},
)
}
#[tokio::test]
async fn usage_shutdown_waits_for_a_producer_before_closing_admission() {
let runtime = UsageRuntime::new(config()).unwrap();
let store = store();
let producer = runtime.track_producer();
let copy = runtime.clone();
let task = tokio::spawn(async move { copy.shutdown(Duration::from_secs(3)).await });
sleep(Duration::from_millis(30)).await;
assert!(!task.is_finished());
runtime
.submit_terminal_event(&store, terminal("last-producer"))
.await;
drop(producer);
task.await.unwrap().unwrap();
assert_eq!(runtime.local_work_pending(), 0);
assert_eq!(
store
.queue
.stats(&runtime.config.stream_key, None)
.await
.unwrap()
.stream_length,
1
);
runtime.shutdown(Duration::from_secs(1)).await.unwrap();
runtime
.submit_terminal_event(&store, terminal("closed"))
.await;
runtime
.record_terminal_event(&store, terminal("closed-direct-api"))
.await;
assert_eq!(
store
.queue
.stats(&runtime.config.stream_key, None)
.await
.unwrap()
.stream_length,
1
);
}
#[tokio::test]
async fn usage_shutdown_persists_concurrent_terminal_handoffs() {
let runtime = UsageRuntime::new(config()).unwrap();
let store = store();
let mut tasks = tokio::task::JoinSet::new();
for index in 0..128 {
let runtime = runtime.clone();
let store = store.clone();
let producer = runtime.track_producer();
tasks.spawn(async move {
let _producer = producer;
runtime
.submit_terminal_event(&store, terminal(&format!("drain-{index}")))
.await;
});
}
runtime.shutdown(Duration::from_secs(5)).await.unwrap();
while let Some(result) = tasks.join_next().await {
result.unwrap();
}
assert_eq!(
store
.queue
.stats(&runtime.config.stream_key, None)
.await
.unwrap()
.stream_length,
128
);
assert_eq!(runtime.local_work_pending(), 0);
}
#[tokio::test]
async fn usage_shutdown_wakes_retry_backoff_and_preserves_all_buffered_events() {
let runtime = UsageRuntime::new(config()).unwrap();
let store = store();
let queue = UsageQueue::new(Arc::clone(&store.queue), runtime.config.clone()).unwrap();
for index in 0..32 {
assert!(runtime.enqueue_retry.schedule(
queue.clone(),
terminal(&format!("retry-{index}")),
"terminal",
DataLayerError::Redis("transient".into())
));
}
runtime.shutdown(Duration::from_secs(3)).await.unwrap();
assert_eq!(runtime.enqueue_retry.pending(), 0);
assert_eq!(runtime.enqueue_retry.recovered_total(), 32);
assert_eq!(
store
.queue
.stats(&runtime.config.stream_key, None)
.await
.unwrap()
.stream_length,
32
);
}
#[tokio::test]
async fn usage_shutdown_failure_retains_retry_work_for_a_later_attempt() {
let runtime = UsageRuntime::new(config()).unwrap();
let store = store();
let flaky = Arc::new(FlakyAppendQueueStore::new(
Arc::clone(&store.queue),
usize::MAX,
));
let queue = UsageQueue::new(flaky.clone(), runtime.config.clone()).unwrap();
assert!(runtime.enqueue_retry.schedule(
queue,
terminal("recover-after-deadline"),
"terminal",
DataLayerError::Redis("unavailable".into())
));
let result = runtime.shutdown(Duration::from_millis(150)).await;
assert!(matches!(result, Err(DataLayerError::TimedOut(_))));
assert_eq!(runtime.enqueue_retry.pending(), 1);
assert!(flaky.append_attempts.load(Ordering::Acquire) <= 4);
flaky.remaining_failures.store(0, Ordering::Release);
runtime.shutdown(Duration::from_secs(3)).await.unwrap();
assert_eq!(runtime.enqueue_retry.recovered_total(), 1);
assert_eq!(
store
.queue
.stats(&runtime.config.stream_key, None)
.await
.unwrap()
.stream_length,
1
);
}
#[tokio::test]
async fn usage_shutdown_flushes_delayed_lifecycle_without_waiting_for_its_timer() {
let mut config = config();
config.lifecycle_enqueue_delay_ms = 60_000;
let runtime = UsageRuntime::new(config).unwrap();
let store = store();
let event = UsageEvent::new(
UsageEventType::Pending,
"delayed",
UsageEventData::default(),
);
runtime
.enqueue_lifecycle_event_with_config_delay(&store, event)
.await;
assert!(runtime.local_work_pending() > 0);
runtime.shutdown(Duration::from_secs(3)).await.unwrap();
assert_eq!(runtime.local_work_pending(), 0);
assert_eq!(
store
.queue
.stats(&runtime.config.stream_key, None)
.await
.unwrap()
.stream_length,
1
);
}
#[tokio::test]
async fn usage_shutdown_stops_every_idle_worker_and_supervisor() {
for supervised in [false, true] {
let runtime = UsageRuntime::new(config()).unwrap();
let store = Arc::new(store());
let handles = if supervised {
vec![runtime.spawn_worker_supervisor(store).unwrap()]
} else {
runtime.spawn_workers(store)
};
sleep(Duration::from_millis(30)).await;
runtime.shutdown(Duration::from_secs(3)).await.unwrap();
for handle in handles {
timeout(Duration::from_secs(1), handle)
.await
.unwrap()
.unwrap();
}
assert_eq!(runtime.metrics_snapshot().worker_active_count, 0);
assert_eq!(runtime.shutdown.supervisors.load(Ordering::Acquire), 0);
}
}
#[tokio::test]
async fn usage_shutdown_does_not_cancel_a_write_or_ack_its_unfinished_record() {
let runtime = UsageRuntime::new(config()).unwrap();
let store = BlockingWriteQueueConfiguredUsageStore {
records: Arc::new(Mutex::new(Vec::new())),
queue: store().queue,
write_started: Arc::new(tokio::sync::Notify::new()),
release_writes: Arc::new(tokio::sync::Notify::new()),
writes_completed: Arc::new(AtomicUsize::new(0)),
};
let queue = UsageQueue::new(Arc::clone(&store.queue), runtime.config.clone()).unwrap();
queue.ensure_consumer_group().await.unwrap();
queue.enqueue(&terminal("in-flight-worker")).await.unwrap();
let worker = runtime.spawn_worker(Arc::new(store.clone())).unwrap();
timeout(Duration::from_secs(3), store.write_started.notified())
.await
.unwrap();
assert!(runtime.shutdown(Duration::from_millis(50)).await.is_err());
assert!(!worker.is_finished());
assert_eq!(
store
.queue
.stats(
&runtime.config.stream_key,
Some(&runtime.config.consumer_group)
)
.await
.unwrap()
.group_pending,
1
);
store.release_writes.notify_one();
runtime.shutdown(Duration::from_secs(3)).await.unwrap();
worker.await.unwrap();
assert_eq!(store.writes_completed.load(Ordering::Acquire), 1);
assert_eq!(
store
.queue
.stats(
&runtime.config.stream_key,
Some(&runtime.config.consumer_group)
)
.await
.unwrap()
.group_pending,
0
);
}
#[tokio::test]
async fn usage_shutdown_disabled_runtime_is_immediate() {
UsageRuntime::disabled()
.shutdown(Duration::from_secs(1))
.await
.unwrap();
}
#[tokio::test]
async fn usage_shutdown_consumes_process_local_queue_before_stopping_workers() {
let runtime = UsageRuntime::new(config()).unwrap();
let store = store();
let worker = runtime
.spawn_worker_supervisor(Arc::new(store.clone()))
.unwrap();
for index in 0..64 {
runtime
.submit_terminal_event(&store, terminal(&format!("local-{index}")))
.await;
}
runtime
.shutdown_with_local_queue(Duration::from_secs(5), Some(Arc::clone(&store.queue)))
.await
.unwrap();
worker.await.unwrap();
assert_eq!(store.records.lock().unwrap().len(), 64);
assert_eq!(
store
.queue
.stats(
&runtime.config.stream_key,
Some(&runtime.config.consumer_group)
)
.await
.unwrap()
.group_pending,
0
);
}
#[tokio::test]
async fn usage_shutdown_rejects_unconsumed_memory_queue_as_success() {
let runtime = UsageRuntime::new(config()).unwrap();
let store = store();
runtime
.submit_terminal_event(&store, terminal("unconsumed"))
.await;
let result = runtime
.shutdown_with_local_queue(Duration::from_millis(50), Some(Arc::clone(&store.queue)))
.await;
assert!(result.is_err());
let worker = runtime.spawn_worker(Arc::new(store.clone())).unwrap();
runtime
.shutdown_with_local_queue(Duration::from_secs(3), Some(Arc::clone(&store.queue)))
.await
.unwrap();
worker.await.unwrap();
assert_eq!(store.records.lock().unwrap().len(), 1);
}
#[tokio::test]
async fn usage_shutdown_does_not_miss_accepted_pending_to_terminal_handoffs() {
let mut config = config();
config.queue_terminal_events = false;
let runtime = UsageRuntime::new(config).unwrap();
let store = store();
for index in 0..32 {
let id = format!("ordered-{index}");
let plan = terminal_test_plan(&id);
runtime.record_pending(&store, build_lifecycle_usage_seed(&plan, None));
runtime.record_stream_started(
&store,
&build_lifecycle_usage_seed(&plan, None),
200,
Some(&ExecutionTelemetry {
ttfb_ms: Some(5),
elapsed_ms: None,
upstream_bytes: None,
}),
);
runtime.submit_terminal_event(&store, terminal(&id)).await;
}
runtime.shutdown(Duration::from_secs(5)).await.unwrap();
let records = store.records.lock().unwrap();
for index in 0..32 {
let statuses: Vec<_> = records
.iter()
.filter(|r| r.request_id == format!("ordered-{index}"))
.map(|r| r.status.as_str())
.collect();
assert_eq!(statuses, ["pending", "streaming", "completed"]);
}
}
+63 -27
View File
@@ -36,8 +36,19 @@ pub async fn reconcile_usage_policy_cost_for_event(
writer: &dyn UsageSettlementWriter,
event: &UsageEvent,
) -> Result<(), DataLayerError> {
reconcile_usage_policy_cost_for_event_with_result(writer, event)
.await
.map(|_| ())
}
pub(crate) struct ReconciledUsagePolicyCost(ReconcileUsagePolicyCostInput);
pub(crate) async fn reconcile_usage_policy_cost_for_event_with_result(
writer: &dyn UsageSettlementWriter,
event: &UsageEvent,
) -> Result<Option<ReconciledUsagePolicyCost>, DataLayerError> {
if !writer.has_usage_settlement_writer() {
return Ok(());
return Ok(None);
}
let terminal_state = match event.event_type {
UsageEventType::Completed => UsagePolicyCostReservationState::Finalized,
@@ -49,16 +60,16 @@ pub async fn reconcile_usage_policy_cost_for_event(
UsageEventType::Failed | UsageEventType::Cancelled => {
UsagePolicyCostReservationState::Released
}
UsageEventType::Pending | UsageEventType::Streaming => return Ok(()),
UsageEventType::Pending | UsageEventType::Streaming => return Ok(None),
};
if plan_usage_reservation_reconciliation_is_deferred(event.data.request_metadata.as_ref()) {
return Ok(());
return Ok(None);
}
let Some(subject_id) = event.data.user_id.as_deref().and_then(non_empty_trimmed) else {
return Ok(());
return Ok(None);
};
let Some(reservation_token) = event_usage_policy_reservation_token(event) else {
return Ok(());
return Ok(None);
};
let actual_cost_units = if terminal_state == UsagePolicyCostReservationState::Finalized {
let actual_cost_usd = event.data.actual_total_cost_usd.ok_or_else(|| {
@@ -75,22 +86,39 @@ pub async fn reconcile_usage_policy_cost_for_event(
0
};
let _ = writer
.reconcile_usage_policy_cost(ReconcileUsagePolicyCostInput {
request_id: event.request_id.clone(),
subject_id: subject_id.to_string(),
reservation_token: reservation_token.to_string(),
actual_cost_units,
terminal_state,
finalized_at_unix_secs: event.timestamp_ms / 1_000,
})
.await?;
Ok(())
let input = ReconcileUsagePolicyCostInput {
request_id: event.request_id.clone(),
subject_id: subject_id.to_string(),
reservation_token: reservation_token.to_string(),
actual_cost_units,
terminal_state,
finalized_at_unix_secs: event.timestamp_ms / 1_000,
};
let stored = writer.reconcile_usage_policy_cost(input.clone()).await?;
// A successful call alone is insufficient: None or a different existing terminal
// reservation must not suppress reconciliation of the subsequent stored usage row.
let matches = stored.is_some_and(|stored| {
stored.request_id == input.request_id
&& stored.subject_id == input.subject_id
&& stored.reservation_token == input.reservation_token
&& stored.actual_cost_units == Some(input.actual_cost_units)
&& stored.state == input.terminal_state
&& stored.finalized_at_unix_secs == Some(input.finalized_at_unix_secs)
});
Ok(matches.then_some(ReconciledUsagePolicyCost(input)))
}
pub async fn settle_usage_if_needed(
writer: &dyn UsageSettlementWriter,
usage: &StoredRequestUsageAudit,
) -> Result<(), DataLayerError> {
settle_usage_with_reconciled_cost(writer, usage, None).await
}
pub(crate) async fn settle_usage_with_reconciled_cost(
writer: &dyn UsageSettlementWriter,
usage: &StoredRequestUsageAudit,
reconciled: Option<ReconciledUsagePolicyCost>,
) -> Result<(), DataLayerError> {
if !writer.has_usage_settlement_writer() {
return Ok(());
@@ -132,17 +160,21 @@ pub async fn settle_usage_if_needed(
} else {
(UsagePolicyCostReservationState::Released, 0)
};
let _ = writer
.reconcile_usage_policy_cost(ReconcileUsagePolicyCostInput {
request_id: usage.request_id.clone(),
subject_id: subject_id.to_string(),
reservation_token: reservation_token.to_string(),
actual_cost_units,
terminal_state,
finalized_at_unix_secs: finalized_at_unix_secs
.unwrap_or(usage.updated_at_unix_secs),
})
.await?;
let input = ReconcileUsagePolicyCostInput {
request_id: usage.request_id.clone(),
subject_id: subject_id.to_string(),
reservation_token: reservation_token.to_string(),
actual_cost_units,
terminal_state,
finalized_at_unix_secs: finalized_at_unix_secs
.unwrap_or(usage.updated_at_unix_secs),
};
if !reconciled
.as_ref()
.is_some_and(|previous| previous.0 == input)
{
let _ = writer.reconcile_usage_policy_cost(input).await?;
}
}
}
@@ -243,6 +275,10 @@ fn finite_cost(value: f64) -> Result<f64, DataLayerError> {
#[cfg(test)]
mod tests {
mod reconciliation_reuse {
include!("settlement_reuse_tests.rs");
}
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Mutex;
use std::time::Duration;
@@ -0,0 +1,517 @@
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use aether_data::repository::settlement::InMemorySettlementRepository;
use aether_data_contracts::repository::settlement::{
ReconcileUsagePolicyCostInput, StoredUsagePolicyCostReservation, StoredUsageSettlement,
UsagePolicyCostReservationState,
};
use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UpsertUsageRecord};
use aether_data_contracts::DataLayerError;
use aether_runtime_state::RuntimeQueueStore;
use async_trait::async_trait;
use serde_json::json;
use super::sample_usage as base_usage;
use crate::settlement::{UsageSettlementInput, UsageSettlementWriter};
use crate::worker::write_event_record;
use crate::{
UsageBillingEventEnricher, UsageEvent, UsageEventData, UsageEventType, UsageRecordWriter,
UsageRuntime, UsageRuntimeAccess, UsageRuntimeConfig,
};
const RESERVATION_TOKEN: &str = "550e8400-e29b-41d4-a716-446655440000";
fn sample_usage() -> StoredRequestUsageAudit {
let mut usage = base_usage();
usage.request_metadata = Some(json!({"plan_usage_reservation_token": RESERVATION_TOKEN}));
usage
}
fn runtime() -> UsageRuntime {
UsageRuntime::new(UsageRuntimeConfig {
enabled: true,
..Default::default()
})
.unwrap()
}
#[derive(Default)]
enum ReconcileResponse {
#[default]
Exact,
Missing,
Changed(fn(&mut StoredUsagePolicyCostReservation)),
Error,
}
#[derive(Default)]
struct ReuseStore {
repository: Option<InMemorySettlementRepository>,
response: ReconcileResponse,
stored_override: Option<StoredRequestUsageAudit>,
fail_next_upsert: AtomicBool,
upserts: AtomicUsize,
reconciliations: Mutex<Vec<ReconcileUsagePolicyCostInput>>,
settlements: Mutex<Vec<UsageSettlementInput>>,
}
#[async_trait]
impl UsageSettlementWriter for ReuseStore {
fn has_usage_settlement_writer(&self) -> bool {
true
}
async fn reconcile_usage_policy_cost(
&self,
input: ReconcileUsagePolicyCostInput,
) -> Result<Option<StoredUsagePolicyCostReservation>, DataLayerError> {
input.validate()?;
self.reconciliations.lock().unwrap().push(input.clone());
tokio::task::yield_now().await;
if let Some(repository) = self.repository.as_ref() {
return aether_data_contracts::repository::settlement::SettlementWriteRepository::reconcile_usage_policy_cost(repository, input).await;
}
let mut stored = StoredUsagePolicyCostReservation {
request_id: input.request_id,
subject_id: input.subject_id,
reservation_token: input.reservation_token,
admitted_at_unix_secs: 100,
reserved_cost_units: 100_000_000,
actual_cost_units: Some(input.actual_cost_units),
state: input.terminal_state,
reservation_expires_at_unix_secs: 500,
retain_until_unix_secs: 1_000,
finalized_at_unix_secs: Some(input.finalized_at_unix_secs),
};
match self.response {
ReconcileResponse::Exact => {}
ReconcileResponse::Missing => return Ok(None),
ReconcileResponse::Changed(change) => change(&mut stored),
ReconcileResponse::Error => {
return Err(DataLayerError::TimedOut("reconciliation".to_string()));
}
}
Ok(Some(stored))
}
async fn settle_usage(
&self,
input: UsageSettlementInput,
) -> Result<Option<StoredUsageSettlement>, DataLayerError> {
self.settlements.lock().unwrap().push(input.clone());
if let Some(repository) = self.repository.as_ref() {
return aether_data_contracts::repository::settlement::SettlementWriteRepository::settle_usage(repository, input).await;
}
Ok(None)
}
}
#[async_trait]
impl UsageRecordWriter for ReuseStore {
async fn upsert_usage_record(
&self,
record: UpsertUsageRecord,
) -> Result<Option<StoredRequestUsageAudit>, DataLayerError> {
self.upserts.fetch_add(1, Ordering::Relaxed);
if self.fail_next_upsert.swap(false, Ordering::Relaxed) {
return Err(DataLayerError::TimedOut("upsert".to_string()));
}
if let Some(stored) = self.stored_override.as_ref() {
return Ok(Some(stored.clone()));
}
let mut stored = sample_usage();
stored.request_id = record.request_id;
stored.user_id = record.user_id;
stored.api_key_id = record.api_key_id;
stored.provider_id = record.provider_id;
stored.status = record.status;
stored.billing_status = record.billing_status;
stored.total_cost_usd = record.total_cost_usd.unwrap_or_default();
stored.actual_total_cost_usd = record.actual_total_cost_usd.unwrap_or_default();
stored.request_metadata = record.request_metadata;
stored.updated_at_unix_secs = record.updated_at_unix_secs;
stored.finalized_at_unix_secs = record.finalized_at_unix_secs;
Ok(Some(stored))
}
}
#[async_trait]
impl UsageBillingEventEnricher for ReuseStore {
async fn enrich_usage_event(&self, _event: &mut UsageEvent) -> Result<(), DataLayerError> {
Ok(())
}
}
impl UsageRuntimeAccess for ReuseStore {
fn has_usage_writer(&self) -> bool {
true
}
fn has_usage_worker_queue(&self) -> bool {
false
}
fn usage_worker_queue(&self) -> Option<Arc<dyn RuntimeQueueStore>> {
None
}
}
fn event() -> UsageEvent {
let mut event = UsageEvent::new(
UsageEventType::Completed,
"req-1",
UsageEventData {
user_id: Some("user-1".to_string()),
api_key_id: Some("key-1".to_string()),
provider_name: "openai".to_string(),
model: "gpt-5".to_string(),
total_cost_usd: Some(1.25),
actual_total_cost_usd: Some(0.75),
request_metadata: Some(json!({"plan_usage_reservation_token": RESERVATION_TOKEN})),
..Default::default()
},
);
event.timestamp_ms = 200_999;
event
}
async fn write(store: &ReuseStore, event: UsageEvent, direct: bool) {
if direct {
runtime().record_terminal_event_direct(store, event).await;
} else {
write_event_record(store, &event).await.unwrap();
}
}
#[tokio::test]
async fn worker_and_direct_writes_reuse_confirmed_reservation_and_still_settle_wallet() {
for direct in [false, true] {
let store = ReuseStore::default();
write(&store, event(), direct).await;
assert_eq!(store.upserts.load(Ordering::Relaxed), 1);
let reconciliations = store.reconciliations.lock().unwrap();
assert_eq!(reconciliations.len(), 1, "direct={direct}");
assert_eq!(reconciliations[0].actual_cost_units, 75_000_000);
assert_eq!(reconciliations[0].finalized_at_unix_secs, 200);
let settlements = store.settlements.lock().unwrap();
assert_eq!(settlements.len(), 1);
assert_eq!(settlements[0].request_id, "req-1");
assert_eq!(settlements[0].actual_total_cost_usd, 0.75);
}
}
#[tokio::test]
async fn missing_or_different_reconciliation_results_keep_stored_usage_reconciliation() {
let changes: [fn(&mut StoredUsagePolicyCostReservation); 9] = [
|row| row.request_id = "other-request".to_string(),
|row| row.subject_id = "other-user".to_string(),
|row| row.reservation_token = "other-token".to_string(),
|row| row.actual_cost_units = Some(1),
|row| row.actual_cost_units = None,
|row| row.state = UsagePolicyCostReservationState::Reserved,
|row| row.state = UsagePolicyCostReservationState::Released,
|row| row.finalized_at_unix_secs = Some(199),
|row| row.finalized_at_unix_secs = None,
];
for direct in [false, true] {
for response in std::iter::once(ReconcileResponse::Missing)
.chain(changes.into_iter().map(ReconcileResponse::Changed))
{
let store = ReuseStore {
response,
..Default::default()
};
write(&store, event(), direct).await;
assert_eq!(store.reconciliations.lock().unwrap().len(), 2);
assert_eq!(store.settlements.lock().unwrap().len(), 1);
}
}
}
#[tokio::test]
async fn changed_stored_usage_is_reconciled_using_its_own_identity_cost_and_terminal_state() {
let changes: [fn(&mut StoredRequestUsageAudit); 6] = [
|row| row.request_id = "other-request".to_string(),
|row| row.user_id = Some("other-user".to_string()),
|row| {
row.request_metadata.as_mut().unwrap()["plan_usage_reservation_token"] =
json!("other-token")
},
|row| row.actual_total_cost_usd = 0.25,
|row| row.status = "failed".to_string(),
|row| row.finalized_at_unix_secs = Some(199),
];
for direct in [false, true] {
for change in changes {
let mut stored = sample_usage();
change(&mut stored);
let store = ReuseStore {
stored_override: Some(stored.clone()),
..Default::default()
};
write(&store, event(), direct).await;
let reconciliations = store.reconciliations.lock().unwrap();
assert_eq!(reconciliations.len(), 2);
assert_eq!(reconciliations[1].request_id, stored.request_id);
assert_eq!(
reconciliations[1].subject_id,
stored.user_id.as_ref().unwrap().as_str()
);
assert_eq!(
reconciliations[1].reservation_token,
stored.request_metadata.as_ref().unwrap()["plan_usage_reservation_token"]
.as_str()
.unwrap()
);
assert_ne!(reconciliations[0], reconciliations[1]);
let settlements = store.settlements.lock().unwrap();
assert_eq!(settlements.len(), 1);
assert_eq!(
settlements[0].actual_total_cost_usd,
stored.actual_total_cost_usd
);
assert_eq!(settlements[0].status, stored.status);
}
}
}
#[tokio::test]
async fn cancellation_release_billable_cancellation_and_zero_cost_preserve_settlement_rules() {
for direct in [false, true] {
for (event_type, billable_cancel, cost, terminal_state, wallets) in [
(
UsageEventType::Failed,
false,
0.0,
UsagePolicyCostReservationState::Released,
0,
),
(
UsageEventType::Cancelled,
false,
0.75,
UsagePolicyCostReservationState::Released,
0,
),
(
UsageEventType::Cancelled,
true,
0.75,
UsagePolicyCostReservationState::Finalized,
1,
),
(
UsageEventType::Completed,
false,
0.0,
UsagePolicyCostReservationState::Finalized,
1,
),
] {
let store = ReuseStore::default();
let mut event = event();
event.event_type = event_type;
event.data.actual_total_cost_usd = Some(cost);
event.data.request_metadata.as_mut().unwrap()["cancelled_request_fee"] =
json!(billable_cancel);
write(&store, event, direct).await;
let reconciliations = store.reconciliations.lock().unwrap();
assert_eq!(reconciliations.len(), 1);
assert_eq!(reconciliations[0].terminal_state, terminal_state);
assert_eq!(
reconciliations[0].actual_cost_units,
if billable_cancel { 75_000_000 } else { 0 }
);
assert_eq!(store.settlements.lock().unwrap().len(), wallets);
}
}
}
#[tokio::test]
async fn reconciliation_failure_stops_both_writes_before_upsert_and_wallet_settlement() {
for direct in [false, true] {
let store = ReuseStore {
response: ReconcileResponse::Error,
..Default::default()
};
if direct {
write(&store, event(), true).await;
} else {
assert!(write_event_record(&store, &event()).await.is_err());
}
assert_eq!(store.reconciliations.lock().unwrap().len(), 1);
assert_eq!(store.upserts.load(Ordering::Relaxed), 0);
assert!(store.settlements.lock().unwrap().is_empty());
}
}
#[tokio::test]
async fn retry_after_upsert_failure_reconciles_again_before_settling() {
for direct in [false, true] {
let store = ReuseStore {
fail_next_upsert: AtomicBool::new(true),
..Default::default()
};
if direct {
write(&store, event(), true).await;
} else {
assert!(write_event_record(&store, &event()).await.is_err());
}
assert!(store.settlements.lock().unwrap().is_empty());
write(&store, event(), direct).await;
assert_eq!(store.reconciliations.lock().unwrap().len(), 2);
assert_eq!(store.upserts.load(Ordering::Relaxed), 2);
assert_eq!(store.settlements.lock().unwrap().len(), 1);
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrent_same_user_worker_and_direct_writes_each_reconcile_once() {
const REQUESTS: usize = 256;
let store = Arc::new(ReuseStore::default());
let runtime = runtime();
let barrier = Arc::new(tokio::sync::Barrier::new(REQUESTS));
let mut tasks = tokio::task::JoinSet::new();
for index in 0..REQUESTS {
let store = store.clone();
let runtime = runtime.clone();
let barrier = barrier.clone();
tasks.spawn(async move {
let mut event = event();
event.request_id = format!("reuse-concurrent-{index}");
event.data.request_metadata.as_mut().unwrap()["plan_usage_reservation_token"] =
json!(format!("550e8400-e29b-41d4-a716-{index:012x}"));
barrier.wait().await;
if index % 2 == 0 {
runtime
.record_terminal_event_direct(store.as_ref(), event)
.await;
} else {
write_event_record(store.as_ref(), &event).await.unwrap();
}
});
}
tokio::time::timeout(std::time::Duration::from_secs(10), async {
while let Some(result) = tasks.join_next().await {
result.unwrap();
}
})
.await
.unwrap();
let reconciliations = store.reconciliations.lock().unwrap();
assert_eq!(reconciliations.len(), REQUESTS);
let unique_tokens: std::collections::HashSet<_> = reconciliations
.iter()
.map(|input| &input.reservation_token)
.collect();
assert_eq!(unique_tokens.len(), REQUESTS);
assert_eq!(store.upserts.load(Ordering::Relaxed), REQUESTS);
let settlements = store.settlements.lock().unwrap();
assert_eq!(settlements.len(), REQUESTS);
let unique_requests: std::collections::HashSet<_> =
settlements.iter().map(|input| &input.request_id).collect();
assert_eq!(unique_requests.len(), REQUESTS);
assert_eq!(
settlements
.iter()
.map(|input| input.actual_total_cost_usd)
.sum::<f64>(),
REQUESTS as f64 * 0.75
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrent_duplicate_delivery_debits_real_memory_wallet_only_once() {
use aether_data::repository::wallet::{
InMemoryWalletRepository, StoredWalletSnapshot, WalletLookupKey, WalletReadRepository,
};
use aether_data_contracts::repository::settlement::{
ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, SettlementWriteRepository,
UsagePolicyCostWindow,
};
let wallet = StoredWalletSnapshot::new(
"wallet-1".to_string(),
Some("user-1".to_string()),
None,
10.0,
2.0,
"finite".to_string(),
"USD".to_string(),
"active".to_string(),
0.0,
0.0,
0.0,
0.0,
100,
)
.unwrap();
let wallets = Arc::new(InMemoryWalletRepository::seed([wallet]));
let repository = InMemorySettlementRepository::from_wallet_repository(wallets.clone());
let reservation = ReserveUsagePolicyCostInput {
request_id: "req-1".to_string(),
subject_id: "user-1".to_string(),
reservation_token: RESERVATION_TOKEN.to_string(),
admitted_at_unix_secs: 100,
reserved_cost_units: 100_000_000,
reservation_expires_at_unix_secs: 500,
retain_until_unix_secs: 1_000,
windows: vec![UsagePolicyCostWindow {
window_id: "window-1".to_string(),
starts_at_unix_secs: 0,
ends_at_unix_secs: 1_000,
limit_cost_units: 1_000_000_000,
}],
};
assert!(matches!(
repository
.reserve_usage_policy_cost(reservation.clone())
.await
.unwrap(),
ReserveUsagePolicyCostOutcome::Allowed { .. }
));
let store = Arc::new(ReuseStore {
repository: Some(repository),
..Default::default()
});
let runtime = runtime();
let mut tasks = tokio::task::JoinSet::new();
for index in 0..32 {
let store = store.clone();
let runtime = runtime.clone();
tasks.spawn(async move {
if index % 2 == 0 {
write_event_record(store.as_ref(), &event()).await.unwrap();
} else {
runtime
.record_terminal_event_direct(store.as_ref(), event())
.await;
}
});
}
while let Some(result) = tasks.join_next().await {
result.unwrap();
}
assert_eq!(store.reconciliations.lock().unwrap().len(), 32);
assert_eq!(store.settlements.lock().unwrap().len(), 32);
let wallet = wallets
.find(WalletLookupKey::UserId("user-1"))
.await
.unwrap()
.unwrap();
assert_eq!(wallet.balance + wallet.gift_balance, 11.25);
assert_eq!(wallet.total_consumed, 0.75);
assert!(matches!(
store
.repository
.as_ref()
.unwrap()
.reserve_usage_policy_cost(reservation)
.await
.unwrap(),
ReserveUsagePolicyCostOutcome::AlreadyTerminal {
state: UsagePolicyCostReservationState::Finalized
}
));
}
@@ -0,0 +1,99 @@
use std::future::Future;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use tokio::sync::watch;
use tokio::task::JoinHandle;
#[derive(Debug, Default)]
pub(crate) struct UsageBackgroundTasks {
handles: Mutex<Vec<JoinHandle<()>>>,
}
impl UsageBackgroundTasks {
pub(crate) fn spawn(&self, task: impl Future<Output = ()> + Send + 'static) {
self.handles
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.push(crate::executor::spawn_on_usage_background_runtime(task));
}
pub(crate) async fn stop_idle(&self) {
let handles = std::mem::take(
&mut *self
.handles
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner()),
);
for handle in &handles {
handle.abort();
}
for handle in handles {
let _ = handle.await;
}
}
}
impl Drop for UsageBackgroundTasks {
fn drop(&mut self) {
for handle in self.handles.get_mut().unwrap_or_else(|p| p.into_inner()) {
handle.abort();
}
}
}
#[derive(Debug)]
pub(crate) struct UsageShutdownState {
pub(crate) producers: Arc<AtomicUsize>,
pub(crate) drain: watch::Sender<bool>,
pub(crate) tasks: UsageBackgroundTasks,
pub(crate) worker_control: crate::worker::UsageWorkerControl,
pub(crate) supervisors: Arc<AtomicUsize>,
pub(crate) lock: tokio::sync::Mutex<()>,
}
impl Default for UsageShutdownState {
fn default() -> Self {
Self {
producers: Arc::new(AtomicUsize::new(0)),
drain: watch::channel(false).0,
tasks: UsageBackgroundTasks::default(),
worker_control: crate::worker::UsageWorkerControl::default(),
supervisors: Arc::new(AtomicUsize::new(0)),
lock: tokio::sync::Mutex::new(()),
}
}
}
/// Retain across an owned request finalizer, including any detached handoff.
#[derive(Debug)]
pub struct UsageProducerGuard(pub(crate) Arc<AtomicUsize>);
impl Drop for UsageProducerGuard {
fn drop(&mut self) {
self.0.fetch_sub(1, Ordering::AcqRel);
}
}
pub(crate) async fn wait_for_drain(signal: &mut watch::Receiver<bool>) {
loop {
if *signal.borrow_and_update() {
return;
}
if signal.changed().await.is_err() {
std::future::pending::<()>().await;
}
}
}
pub(crate) async fn retry_delay(delay: Duration, signal: &mut watch::Receiver<bool>) {
if *signal.borrow_and_update() {
tokio::time::sleep(delay.min(Duration::from_millis(100))).await;
} else {
tokio::select! {
_ = tokio::time::sleep(delay) => {},
_ = wait_for_drain(signal) => {},
}
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,449 @@
use super::*;
use crate::dead_letter_encoding::DeadLetterEncodingBudget;
use crate::worker::UsageWorkerObservation;
use aether_runtime_state::RuntimeQueueTransferOutcome;
struct TransferProbe {
inner: RuntimeState,
native: bool,
fail_write: AtomicBool,
lose_reply: AtomicBool,
ack_calls: Mutex<Vec<Vec<String>>>,
append_calls: AtomicUsize,
}
impl TransferProbe {
fn new(native: bool) -> Self {
Self {
inner: RuntimeState::memory(MemoryRuntimeStateConfig::default()),
native,
fail_write: AtomicBool::new(false),
lose_reply: AtomicBool::new(false),
ack_calls: Mutex::new(Vec::new()),
append_calls: AtomicUsize::new(0),
}
}
}
#[async_trait]
impl RuntimeQueueStore for TransferProbe {
async fn ensure_consumer_group(
&self,
stream: &str,
group: &str,
start_id: &str,
) -> Result<(), DataLayerError> {
self.inner
.ensure_consumer_group(stream, group, start_id)
.await
}
async fn append_fields_with_maxlen(
&self,
stream: &str,
fields: &BTreeMap<String, String>,
maxlen: Option<usize>,
) -> Result<String, DataLayerError> {
self.append_calls.fetch_add(1, Ordering::Relaxed);
if self.fail_write.load(Ordering::Acquire) {
return Err(DataLayerError::TimedOut("test append failure".to_string()));
}
self.inner
.append_fields_with_maxlen(stream, fields, maxlen)
.await
}
async fn try_transfer_pending_to_stream(
&self,
source: &str,
group: &str,
entry_id: &str,
destination: &str,
fields: &BTreeMap<String, String>,
) -> Result<Option<RuntimeQueueTransferOutcome>, DataLayerError> {
if !self.native {
return Ok(None);
}
if self.fail_write.load(Ordering::Acquire) {
return Err(DataLayerError::TimedOut(
"test atomic transfer failure".to_string(),
));
}
let result = self
.inner
.try_transfer_pending_to_stream(source, group, entry_id, destination, fields)
.await?;
if self.lose_reply.swap(false, Ordering::AcqRel) {
return Err(DataLayerError::TimedOut(
"test committed transfer reply lost".to_string(),
));
}
Ok(result)
}
async fn read_group(
&self,
stream: &str,
group: &str,
consumer: &str,
count: usize,
block_ms: Option<u64>,
) -> Result<Vec<RuntimeQueueEntry>, DataLayerError> {
self.inner
.read_group(stream, group, consumer, count, block_ms)
.await
}
async fn claim_stale(
&self,
stream: &str,
group: &str,
consumer: &str,
start_id: &str,
config: RuntimeQueueReclaimConfig,
) -> Result<Vec<RuntimeQueueEntry>, DataLayerError> {
self.inner
.claim_stale(stream, group, consumer, start_id, config)
.await
}
async fn ack(
&self,
stream: &str,
group: &str,
ids: &[String],
) -> Result<usize, DataLayerError> {
self.ack_calls.lock().expect("ack calls").push(ids.to_vec());
self.inner.ack(stream, group, ids).await
}
async fn delete(&self, stream: &str, ids: &[String]) -> Result<usize, DataLayerError> {
self.inner.delete(stream, ids).await
}
async fn stats(
&self,
stream: &str,
group: Option<&str>,
) -> Result<RuntimeQueueStats, DataLayerError> {
self.inner.stats(stream, group).await
}
}
async fn transfer_worker(
native: bool,
) -> (
Arc<TransferProbe>,
UsageQueueWorker,
Arc<SelectiveFailingRecorder>,
tokio::sync::mpsc::Receiver<UsageWorkerObservation>,
) {
let runner = Arc::new(TransferProbe::new(native));
let recorder = Arc::new(SelectiveFailingRecorder::default());
let (telemetry, observations) = tokio::sync::mpsc::channel(32);
let mut worker = UsageQueueWorker::new(
runner.clone(),
recorder.clone(),
UsageRuntimeConfig {
enabled: true,
consumer_batch_size: 10,
consumer_block_ms: 1,
..UsageRuntimeConfig::default()
},
None,
)
.expect("worker")
.with_supervisor(UsageWorkerControl::default(), telemetry);
worker.queue =
worker
.queue
.with_dead_letter_encoding_budget(Arc::new(DeadLetterEncodingBudget::new(
64 * 1024 * 1024,
4,
)));
worker.queue.ensure_consumer_group().await.expect("group");
(runner, worker, recorder, observations)
}
async fn malformed_entries(
runner: &TransferProbe,
worker: &UsageQueueWorker,
) -> Vec<RuntimeQueueEntry> {
runner
.inner
.append_fields_with_maxlen(
&worker.config.stream_key,
&BTreeMap::from([
("payload".to_string(), "malformed\u{0000}\n\"\\".to_string()),
(
"legacy".to_string(),
"preserve all original fields".to_string(),
),
]),
None,
)
.await
.expect("raw append");
worker
.queue
.read_group(&worker.consumer)
.await
.expect("read")
}
fn observed_totals(
observations: &mut tokio::sync::mpsc::Receiver<UsageWorkerObservation>,
) -> (usize, usize) {
let mut totals = (0, 0);
while let Ok(observation) = observations.try_recv() {
totals.0 += observation.acked_entries;
totals.1 += observation.dead_lettered_entries;
}
totals
}
#[tokio::test]
async fn native_transfer_replay_does_not_append_or_ack_twice() {
let (runner, worker, recorder, mut observations) = transfer_worker(true).await;
let entries = malformed_entries(&runner, &worker).await;
worker
.process_entries(entries.clone())
.await
.expect("transfer");
worker.process_entries(entries).await.expect("stale replay");
assert_eq!(worker.queue.stats().await.expect("stats").group_pending, 0);
assert_eq!(
worker.queue.dlq_stats().await.expect("dlq").stream_length,
1
);
assert_eq!(observed_totals(&mut observations), (1, 1));
assert!(runner.ack_calls.lock().expect("acks").is_empty());
assert_eq!(runner.append_calls.load(Ordering::Relaxed), 0);
assert!(recorder.calls.lock().expect("record calls").is_empty());
}
#[tokio::test]
async fn committed_transfer_lost_reply_retries_without_duplicate_or_false_metrics() {
let (runner, worker, _, mut observations) = transfer_worker(true).await;
let entries = malformed_entries(&runner, &worker).await;
runner.lose_reply.store(true, Ordering::Release);
assert!(matches!(
worker.process_entries(entries.clone()).await,
Err(DataLayerError::TimedOut(_))
));
worker.process_entries(entries).await.expect("retry");
assert_eq!(
worker.queue.dlq_stats().await.expect("dlq").stream_length,
1
);
assert_eq!(worker.queue.stats().await.expect("stats").stream_length, 0);
assert_eq!(observed_totals(&mut observations), (0, 0));
assert!(runner.ack_calls.lock().expect("acks").is_empty());
assert_eq!(runner.append_calls.load(Ordering::Relaxed), 0);
}
#[tokio::test]
async fn source_no_longer_pending_is_not_reported_as_archived_or_deleted() {
let (runner, worker, _, mut observations) = transfer_worker(true).await;
let entries = malformed_entries(&runner, &worker).await;
runner
.inner
.ack(
&worker.config.stream_key,
&worker.config.consumer_group,
&[entries[0].id.clone()],
)
.await
.expect("external ack");
worker
.process_entries(entries)
.await
.expect("no longer pending");
assert_eq!(worker.queue.stats().await.expect("stats").stream_length, 1);
assert_eq!(
worker.queue.dlq_stats().await.expect("dlq").stream_length,
0
);
assert_eq!(observed_totals(&mut observations), (0, 0));
assert!(runner.ack_calls.lock().expect("acks").is_empty());
}
#[tokio::test]
async fn failed_transfer_acknowledges_successful_prefix_and_preserves_suffix_for_retry() {
let (runner, worker, recorder, mut observations) = transfer_worker(true).await;
for request_id in ["prefix", "req-worker-poison", "suffix"] {
let mut event = sample_event();
event.request_id = request_id.to_string();
worker.queue.enqueue(&event).await.expect("enqueue");
}
let entries = worker
.queue
.read_group(&worker.consumer)
.await
.expect("read batch");
let retry = entries[1..].to_vec();
let prefix_id = entries[0].id.clone();
runner.fail_write.store(true, Ordering::Release);
assert!(worker.process_entries(entries).await.is_err());
assert_eq!(
recorder.calls.lock().expect("calls").as_slice(),
["prefix", "req-worker-poison"]
);
assert_eq!(
*runner.ack_calls.lock().expect("acks"),
vec![vec![prefix_id]]
);
assert_eq!(worker.queue.stats().await.expect("stats").group_pending, 2);
assert_eq!(
worker.queue.dlq_stats().await.expect("dlq").stream_length,
0
);
assert_eq!(observed_totals(&mut observations), (1, 0));
runner.fail_write.store(false, Ordering::Release);
worker.process_entries(retry).await.expect("retry suffix");
assert_eq!(worker.queue.stats().await.expect("stats").stream_length, 0);
assert_eq!(
worker.queue.dlq_stats().await.expect("dlq").stream_length,
1
);
assert_eq!(observed_totals(&mut observations), (2, 1));
assert_eq!(runner.append_calls.load(Ordering::Relaxed), 3);
}
#[tokio::test]
async fn legacy_transfer_fallback_only_acknowledges_after_append_succeeds() {
let (runner, worker, _, mut observations) = transfer_worker(false).await;
let entries = malformed_entries(&runner, &worker).await;
runner.fail_write.store(true, Ordering::Release);
assert!(worker.process_entries(entries.clone()).await.is_err());
assert!(runner.ack_calls.lock().expect("acks").is_empty());
assert_eq!(worker.queue.stats().await.expect("stats").group_pending, 1);
assert_eq!(observed_totals(&mut observations), (0, 0));
runner.fail_write.store(false, Ordering::Release);
worker.process_entries(entries).await.expect("legacy retry");
assert_eq!(
worker.queue.dlq_stats().await.expect("dlq").stream_length,
1
);
assert_eq!(worker.queue.stats().await.expect("stats").stream_length, 0);
assert_eq!(runner.ack_calls.lock().expect("acks").len(), 1);
assert_eq!(observed_totals(&mut observations), (1, 1));
}
#[tokio::test]
async fn encoding_rejection_preserves_pending_original_until_budget_allows_retry() {
let (runner, mut worker, _, mut observations) = transfer_worker(true).await;
let small = Arc::new(DeadLetterEncodingBudget::new(1, 1));
worker.queue = worker.queue.with_dead_letter_encoding_budget(small.clone());
let entries = malformed_entries(&runner, &worker).await;
let original = entries[0].fields.clone();
assert!(worker.process_entries(entries.clone()).await.is_err());
assert_eq!(small.snapshot().reserved_bytes, 0);
assert_eq!(small.snapshot().active_jobs, 0);
assert_eq!(worker.queue.stats().await.expect("stats").group_pending, 1);
assert_eq!(
worker.queue.dlq_stats().await.expect("dlq").stream_length,
0
);
assert_eq!(observed_totals(&mut observations), (0, 0));
worker.queue = worker
.queue
.with_dead_letter_encoding_budget(Arc::new(DeadLetterEncodingBudget::new(64 * 1024, 1)));
worker
.process_entries(entries)
.await
.expect("retry with capacity");
runner
.inner
.ensure_consumer_group(&worker.config.dlq_stream_key, "inspect", "0-0")
.await
.expect("dlq group");
let dlq = runner
.inner
.read_group(
&worker.config.dlq_stream_key,
"inspect",
"inspector",
1,
Some(1),
)
.await
.expect("dlq read");
let payload: serde_json::Value =
serde_json::from_str(&dlq[0].fields["payload"]).expect("wire JSON");
assert_eq!(
payload["fields"],
serde_json::to_value(original).expect("original JSON")
);
assert_eq!(observed_totals(&mut observations), (1, 1));
}
#[tokio::test]
async fn normal_record_replay_reports_actual_ack_count() {
let (_, worker, _, mut observations) = transfer_worker(true).await;
worker
.queue
.enqueue(&sample_event())
.await
.expect("enqueue");
let entries = worker
.queue
.read_group(&worker.consumer)
.await
.expect("read");
worker
.process_entries(entries.clone())
.await
.expect("first record");
worker.process_entries(entries).await.expect("replay");
assert_eq!(observed_totals(&mut observations), (1, 0));
}
#[tokio::test]
async fn oversized_dead_letter_does_not_block_healthy_entries_in_the_same_batch() {
let (runner, mut worker, recorder, mut observations) = transfer_worker(true).await;
let bad = malformed_entries(&runner, &worker).await;
for request_id in ["healthy-first", "healthy-second"] {
let mut event = sample_event();
event.request_id = request_id.to_string();
worker.queue.enqueue(&event).await.expect("healthy enqueue");
}
let mut entries = bad.clone();
entries.extend(
worker
.queue
.read_group(&worker.consumer)
.await
.expect("healthy read"),
);
worker.queue = worker
.queue
.with_dead_letter_encoding_budget(Arc::new(DeadLetterEncodingBudget::new(1, 1)));
assert!(worker.process_entries(entries).await.is_err());
assert_eq!(
recorder.calls.lock().expect("calls").as_slice(),
["healthy-first", "healthy-second"]
);
let stats = worker.queue.stats().await.expect("stats");
assert_eq!((stats.group_pending, stats.stream_length), (1, 1));
assert_eq!(
worker.queue.dlq_stats().await.expect("dlq").stream_length,
0
);
assert_eq!(observed_totals(&mut observations), (2, 0));
assert!(worker.process_entries(bad.clone()).await.is_err());
assert_eq!(observed_totals(&mut observations), (0, 0));
worker.queue = worker
.queue
.with_dead_letter_encoding_budget(Arc::new(DeadLetterEncodingBudget::new(64 * 1024, 1)));
worker
.process_entries(bad)
.await
.expect("archive after raising budget");
assert_eq!(worker.queue.stats().await.expect("stats").group_pending, 0);
assert_eq!(
worker.queue.dlq_stats().await.expect("dlq").stream_length,
1
);
assert_eq!(observed_totals(&mut observations), (1, 1));
}
+3
View File
@@ -581,6 +581,7 @@ fn build_lifecycle_usage_event_from_record(
execution_path: record.execution_path,
local_execution_runtime_miss_reason: record.local_execution_runtime_miss_reason,
request_metadata: record.request_metadata,
capture_retention: record.capture_retention,
..UsageEventData::default()
},
}
@@ -1534,6 +1535,7 @@ fn build_lifecycle_usage_record_owned(
};
Ok(UpsertUsageRecord {
capture_retention: Default::default(),
request_id,
user_id,
api_key_id,
@@ -1631,6 +1633,7 @@ fn build_lifecycle_usage_record_impl(
};
Ok(UpsertUsageRecord {
capture_retention: Default::default(),
request_id: seed.request_id.clone(),
user_id: seed.user_id.clone(),
api_key_id: seed.api_key_id.clone(),