#![cfg(feature = "test-fixtures")]
use klieo_a2a::server::{A2aDispatcher, TaskEvent};
use klieo_a2a::types::TaskStatus;
use klieo_a2a::EchoHandler;
use klieo_auth_common::AllowAnonymous;
use klieo_auth_common::Authenticator;
use klieo_bus_memory::MemoryBus;
use klieo_core::{DurableName, Pubsub};
use opentelemetry::global;
use opentelemetry::trace::TraceContextExt as _;
use opentelemetry_sdk::testing::trace::InMemorySpanExporter;
use opentelemetry_sdk::trace::TracerProvider;
use std::sync::Arc;
use std::time::Duration;
use tokio_stream::StreamExt as _;
use tracing_subscriber::layer::SubscriberExt as _;
use tracing_subscriber::Registry;
#[tokio::test]
async fn publisher_span_parents_subscriber_span_across_replicas() {
let exporter = InMemorySpanExporter::default();
let provider = TracerProvider::builder()
.with_simple_exporter(exporter.clone())
.build();
let tracer = {
use opentelemetry::trace::TracerProvider as _;
provider.tracer("klieo-test")
};
global::set_tracer_provider(provider.clone());
let otel_layer = tracing_opentelemetry::layer().with_tracer(tracer);
let subscriber = Registry::default().with(otel_layer);
let _guard = tracing::subscriber::set_default(subscriber);
let bus = MemoryBus::new();
let pubsub: Arc<dyn Pubsub> = bus.pubsub.clone();
let handler: Arc<dyn klieo_a2a::handler::A2aHandler> = Arc::new(EchoHandler::default());
let auth: Arc<dyn Authenticator> = Arc::new(AllowAnonymous);
let dispatcher = A2aDispatcher::new(handler, auth, pubsub.clone());
let sink = dispatcher.event_sink();
let durable = DurableName::new("trace-stitch-test");
let mut stream = pubsub
.subscribe("klieo.a2a.task.t-trace", durable)
.await
.unwrap();
let test_root_trace_id = {
let root = tracing::info_span!("test_root");
let _enter = root.enter();
let event = TaskEvent::new("t-trace", TaskStatus::Working, None, false).with_event_id(1);
sink.send(event).await.unwrap();
use tracing_opentelemetry::OpenTelemetrySpanExt as _;
root.context().span().span_context().trace_id()
};
let msg = tokio::time::timeout(Duration::from_millis(500), stream.next())
.await
.expect("subscribe timeout")
.expect("stream ended")
.expect("bus error");
assert!(
msg.headers.contains_key("traceparent"),
"publisher must inject W3C traceparent into bus headers; got headers: {:?}",
msg.headers.keys().collect::<Vec<_>>(),
);
let extracted_cx = klieo_core::extract_traceparent(&msg.headers);
let span_ctx = extracted_cx.span().span_context().clone();
assert!(
span_ctx.is_valid(),
"extracted span context must be valid (trace_id={:?}, span_id={:?})",
span_ctx.trace_id(),
span_ctx.span_id(),
);
msg.ack.ack().await.unwrap();
provider.force_flush();
let spans = exporter.get_finished_spans().unwrap();
let publisher_span = spans
.iter()
.find(|s| s.name == "send")
.or_else(|| spans.iter().find(|s| s.name.contains("send")))
.unwrap_or_else(|| {
panic!(
"expected publisher span named `send`; got: {:?}",
spans.iter().map(|s| &s.name).collect::<Vec<_>>(),
)
});
assert_eq!(
publisher_span.span_context.trace_id(),
test_root_trace_id,
"publisher span must inherit test_root's trace_id for cross-replica stitch",
);
}