mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-04 00:17:45 +08:00
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:
@@ -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),
|
||||
|
||||
+282
@@ -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}}),
|
||||
],
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -14,3 +14,6 @@ serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
thiserror.workspace = true
|
||||
tokio.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
aether-runtime-state.workspace = true
|
||||
|
||||
@@ -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()),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
};
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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
@@ -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};
|
||||
|
||||
@@ -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}
|
||||
@@ -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?;
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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"
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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::{
|
||||
|
||||
@@ -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()
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -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"]);
|
||||
}
|
||||
}
|
||||
@@ -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));
|
||||
}
|
||||
@@ -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(),
|
||||
|
||||
Reference in New Issue
Block a user