use crate::{
middleware::get_scope,
util::{http_method_str, http_url},
};
use actix_http::{encoding::Decoder, BoxedPayloadStream, Error, Payload};
use actix_web::{
body::MessageBody,
http::{
self,
header::{HeaderName, HeaderValue},
},
web::Bytes,
};
use awc::{
error::SendRequestError,
http::header::{CONTENT_LENGTH, USER_AGENT},
ClientRequest, ClientResponse,
};
use futures_util::{future::TryFutureExt as _, Future, Stream};
use opentelemetry::{
global,
propagation::Injector,
trace::{SpanKind, Status, TraceContextExt, Tracer},
Context, KeyValue,
};
use opentelemetry_semantic_conventions::{
attribute::MESSAGING_MESSAGE_BODY_SIZE,
trace::{
HTTP_REQUEST_METHOD, HTTP_RESPONSE_STATUS_CODE, SERVER_ADDRESS, SERVER_PORT, URL_FULL,
USER_AGENT_ORIGINAL,
},
};
use serde::Serialize;
use std::mem;
use std::str::FromStr;
use std::{
borrow::Cow,
fmt::{self, Debug},
};
pub struct InstrumentedClientRequest {
cx: Context,
attrs: Vec<KeyValue>,
span_namer: fn(&ClientRequest) -> String,
request: ClientRequest,
}
impl Debug for InstrumentedClientRequest {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let span_namer = fmt::Pointer::fmt(&(self.span_namer as usize as *const ()), f);
f.debug_struct("InstrumentedClientRequest")
.field("cx", &self.cx)
.field("attrs", &self.attrs)
.field("span_namer", &span_namer)
.field("request", &self.request)
.finish()
}
}
fn default_span_namer(request: &ClientRequest) -> String {
format!(
"{} {}",
request.get_method(),
request.get_uri().host().unwrap_or_default()
)
}
pub trait ClientExt {
fn trace_request(self) -> InstrumentedClientRequest
where
Self: Sized,
{
self.trace_request_with_context(Context::current())
}
fn trace_request_with_context(self, cx: Context) -> InstrumentedClientRequest;
}
impl ClientExt for ClientRequest {
fn trace_request_with_context(self, cx: Context) -> InstrumentedClientRequest {
InstrumentedClientRequest {
cx,
attrs: Vec::with_capacity(8),
span_namer: default_span_namer,
request: self,
}
}
}
type AwcResult = Result<ClientResponse<Decoder<Payload<BoxedPayloadStream>>>, SendRequestError>;
impl InstrumentedClientRequest {
pub async fn send(self) -> AwcResult {
self.trace_request(|request| request.send()).await
}
pub async fn send_body<B>(self, body: B) -> AwcResult
where
B: MessageBody + 'static,
{
self.trace_request(|request| request.send_body(body)).await
}
pub async fn send_form<T: Serialize>(self, value: &T) -> AwcResult {
self.trace_request(|request| request.send_form(value)).await
}
pub async fn send_json<T: Serialize>(self, value: &T) -> AwcResult {
self.trace_request(|request| request.send_json(value)).await
}
pub async fn send_stream<S, E>(self, stream: S) -> AwcResult
where
S: Stream<Item = Result<Bytes, E>> + Unpin + 'static,
E: std::error::Error + Into<Error> + 'static,
{
self.trace_request(|request| request.send_stream(stream))
.await
}
async fn trace_request<F, R>(mut self, f: F) -> AwcResult
where
F: FnOnce(ClientRequest) -> R,
R: Future<Output = AwcResult>,
{
let tracer = global::tracer_with_scope(get_scope());
self.attrs.extend(
&mut [
KeyValue::new(
SERVER_ADDRESS,
self.request
.get_uri()
.host()
.map(|u| Cow::Owned(u.to_string()))
.unwrap_or(Cow::Borrowed("unknown")),
),
KeyValue::new(
HTTP_REQUEST_METHOD,
http_method_str(self.request.get_method()),
),
KeyValue::new(URL_FULL, http_url(self.request.get_uri())),
]
.into_iter(),
);
if let Some(peer_port) = self.request.get_uri().port_u16() {
if peer_port != 80 && peer_port != 443 {
self.attrs
.push(KeyValue::new(SERVER_PORT, peer_port as i64));
}
}
if let Some(user_agent) = self
.request
.headers()
.get(USER_AGENT)
.and_then(|len| len.to_str().ok())
{
self.attrs
.push(KeyValue::new(USER_AGENT_ORIGINAL, user_agent.to_string()))
}
if let Some(content_length) = self.request.headers().get(CONTENT_LENGTH).and_then(|len| {
len.to_str()
.ok()
.and_then(|str_len| str_len.parse::<i64>().ok())
}) {
self.attrs
.push(KeyValue::new(MESSAGING_MESSAGE_BODY_SIZE, content_length))
}
let span = tracer
.span_builder((self.span_namer)(&self.request))
.with_kind(SpanKind::Client)
.with_attributes(mem::take(&mut self.attrs))
.start_with_context(&tracer, &self.cx);
let cx = self.cx.with_span(span);
global::get_text_map_propagator(|injector| {
injector.inject_context(&cx, &mut ActixClientCarrier::new(&mut self.request));
});
f(self.request)
.inspect_ok(|res| record_response(res, &cx))
.inspect_err(|err| record_err(err, &cx))
.await
}
pub fn with_attributes(
mut self,
attrs: impl IntoIterator<Item = KeyValue>,
) -> InstrumentedClientRequest {
self.attrs.extend(&mut attrs.into_iter());
self
}
pub fn with_span_namer(
mut self,
span_namer: fn(&ClientRequest) -> String,
) -> InstrumentedClientRequest {
self.span_namer = span_namer;
self
}
}
fn convert_status(status: http::StatusCode) -> Status {
match status.as_u16() {
100..=399 => Status::Unset,
400..=599 => Status::error("Unexpected status code"),
code => Status::error(format!("Invalid HTTP status code {}", code)),
}
}
fn record_response<T>(response: &ClientResponse<T>, cx: &Context) {
let span = cx.span();
let status = convert_status(response.status());
span.set_status(status);
span.set_attribute(KeyValue::new(
HTTP_RESPONSE_STATUS_CODE,
response.status().as_u16() as i64,
));
span.end();
}
fn record_err<T: fmt::Debug>(err: T, cx: &Context) {
let span = cx.span();
span.set_status(Status::error(format!("{:?}", err)));
span.end();
}
struct ActixClientCarrier<'a> {
request: &'a mut ClientRequest,
}
impl<'a> ActixClientCarrier<'a> {
fn new(request: &'a mut ClientRequest) -> Self {
ActixClientCarrier { request }
}
}
impl Injector for ActixClientCarrier<'_> {
fn set(&mut self, key: &str, value: String) {
let header_name = HeaderName::from_str(key).expect("Must be header name");
let header_value = HeaderValue::from_str(&value).expect("Must be a header value");
self.request.headers_mut().insert(header_name, header_value);
}
}