Skip to content

Commit bfb698e

Browse files
committed
fix: keep handoff spans through terminal writes
1 parent 5a18c98 commit bfb698e

2 files changed

Lines changed: 40 additions & 24 deletions

File tree

‎packages/sie_server/src/sie_server/local_ingest_client.py‎

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
import asyncio
66
import hashlib
77
import os
8-
from collections.abc import AsyncIterator
8+
from collections.abc import AsyncGenerator
99
from contextlib import nullcontext, suppress
1010
from typing import Any
1111

@@ -126,13 +126,13 @@ async def _read_frame(reader: asyncio.StreamReader) -> dict[str, Any]:
126126
return response
127127

128128

129-
def stream_generate(socket_path: str, items: bytes, params: bytes, meta: dict[str, Any]) -> AsyncIterator[bytes]:
129+
def stream_generate(socket_path: str, items: bytes, params: bytes, meta: dict[str, Any]) -> AsyncGenerator[bytes, None]:
130130
if os.environ.get("SIE_TRACING_ENABLED", "").strip().lower() not in {"1", "true", "yes", "on"}:
131131
return _stream_generate(socket_path, items, params, meta)
132132
return _traced_generate(socket_path, items, params, meta)
133133

134134

135-
async def _traced_generate(socket_path: str, items: bytes, params: bytes, meta: dict[str, Any]) -> AsyncIterator[bytes]:
135+
async def _traced_generate(socket_path: str, items: bytes, params: bytes, meta: dict[str, Any]) -> AsyncGenerator[bytes, None]:
136136
carrier = {
137137
key: value
138138
for key, limit in (("traceparent", 256), ("tracestate", 512))
@@ -188,7 +188,7 @@ async def _stream_generate(
188188
items: bytes,
189189
params: bytes,
190190
meta: dict[str, Any],
191-
) -> AsyncIterator[bytes]:
191+
) -> AsyncGenerator[bytes, None]:
192192
"""Map one caller operation onto sidecar protocol v0.2.
193193
194194
Pulling one response before requesting the next naturally propagates

‎packages/sie_server_sidecar/src/local_ingest.rs‎

Lines changed: 36 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -929,26 +929,32 @@ async fn handle_connection(stream: UnixStream, shared: Arc<IngestShared>) {
929929
let _generate_claim = generate_claim;
930930
let operation_id = request.id;
931931
let stream_writer = GenerateStreamWriter::new(operation_id, writer_op);
932+
let mut span = tracing::Span::none();
932933
let result = publish_generate_stream(
933934
request.body,
934935
&shared_op,
935936
stream_writer.clone(),
936937
Arc::clone(&lifecycle_op),
938+
&mut span,
937939
)
938940
.await;
939-
active_op.lock().await.remove(&request_id);
940-
if !lifecycle_op.is_closed() {
941-
let write_result = match result {
942-
Ok(()) => stream_writer.finish_success().await,
943-
Err(error) => stream_writer.finish_error(&error).await,
944-
};
945-
if let Err(error) = write_result {
946-
debug!(
947-
request_id,
948-
error, "local-ingest: generation terminal write failed"
949-
);
941+
async {
942+
active_op.lock().await.remove(&request_id);
943+
if !lifecycle_op.is_closed() {
944+
let write_result = match result {
945+
Ok(()) => stream_writer.finish_success().await,
946+
Err(error) => stream_writer.finish_error(&error).await,
947+
};
948+
if let Err(error) = write_result {
949+
debug!(
950+
request_id,
951+
error, "local-ingest: generation terminal write failed"
952+
);
953+
}
950954
}
951955
}
956+
.instrument(span)
957+
.await;
952958
});
953959
continue;
954960
}
@@ -1011,8 +1017,9 @@ async fn handle_connection(stream: UnixStream, shared: Arc<IngestShared>) {
10111017
let _retained_data_permit = retained_data_permit;
10121018
let _inbound_permit = inbound_permit;
10131019
let _operation_id_guard = operation_id_guard;
1014-
let frame = run_op(request, &shared_op).await;
1015-
write_response(&writer_op, frame).await;
1020+
let mut span = tracing::Span::none();
1021+
let frame = run_op(request, &shared_op, &mut span).await;
1022+
write_response(&writer_op, frame).instrument(span).await;
10161023
});
10171024
}
10181025

@@ -1043,10 +1050,14 @@ async fn handle_connection(stream: UnixStream, shared: Arc<IngestShared>) {
10431050
}
10441051
}
10451052

1046-
async fn run_op(request: RequestEnvelope, shared: &IngestShared) -> Vec<u8> {
1053+
async fn run_op(
1054+
request: RequestEnvelope,
1055+
shared: &IngestShared,
1056+
span: &mut tracing::Span,
1057+
) -> Vec<u8> {
10471058
match request.op.as_str() {
10481059
OP_PING => encode_response(request.id, true, None, empty_body()),
1049-
OP_PUBLISH_WORK => match publish_work(request.body, shared).await {
1060+
OP_PUBLISH_WORK => match publish_work(request.body, shared, span).await {
10501061
Ok(results_bytes) => {
10511062
encode_response(request.id, true, None, results_body(results_bytes))
10521063
}
@@ -1150,6 +1161,7 @@ async fn publish_generate_stream(
11501161
shared: &IngestShared,
11511162
stream_writer: GenerateStreamWriter,
11521163
lifecycle: Arc<ConnectionLifecycle>,
1164+
span: &mut tracing::Span,
11531165
) -> Result<(), crate::dispatcher::GenerateDispatchError> {
11541166
let semantic_deadline = generate_timeout_deadline(body.timeout_ms).map_err(|message| {
11551167
crate::dispatcher::GenerateDispatchError {
@@ -1180,7 +1192,7 @@ async fn publish_generate_stream(
11801192
message: "publish_generate_stream requires endpoint generate".to_string(),
11811193
});
11821194
}
1183-
let span = local_ingest_span(&body, &mut items);
1195+
*span = local_ingest_span(&body, &mut items);
11841196
publish_validated_generate(
11851197
body,
11861198
items,
@@ -1190,7 +1202,7 @@ async fn publish_generate_stream(
11901202
semantic_deadline,
11911203
timeout_ms,
11921204
)
1193-
.instrument(span)
1205+
.instrument(span.clone())
11941206
.await
11951207
}
11961208

@@ -1585,15 +1597,19 @@ fn error_result(wi: &WorkItem, worker_id: &str, code: &str, message: &str) -> Wo
15851597
}
15861598
}
15871599

1588-
async fn publish_work(body: RequestBody, shared: &IngestShared) -> Result<Vec<u8>, String> {
1600+
async fn publish_work(
1601+
body: RequestBody,
1602+
shared: &IngestShared,
1603+
span: &mut tracing::Span,
1604+
) -> Result<Vec<u8>, String> {
15891605
validate_publish_work_timeout(body.timeout_ms)?;
15901606
validate_payload_digest(&body)?;
15911607
let mut items: Vec<WorkItem> = rmp_serde::from_slice(&body.items)
15921608
.map_err(|e| format!("DecodeError: items is not a msgpack WorkItem array: {e}"))?;
15931609
validate_work_items(&body, &items)?;
1594-
let span = local_ingest_span(&body, &mut items);
1610+
*span = local_ingest_span(&body, &mut items);
15951611
publish_validated_work(body, items, shared)
1596-
.instrument(span)
1612+
.instrument(span.clone())
15971613
.await
15981614
}
15991615

0 commit comments

Comments
 (0)