use std::fmt;
use std::sync::Arc;
use std::time::{Instant, SystemTime, UNIX_EPOCH};
use http::Request;
use serde::Serialize;
use tokio::net::TcpStream;
use tokio_tungstenite::{MaybeTlsStream, WebSocketStream};
use tungstenite::error::UrlError;
use tungstenite::Error as TungsteniteError;
use uuid::Uuid;
use crate::{DeepgramError, Result};
pub const SCHEMA_VERSION: u32 = 1;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum ConnectOutcome {
Completed,
Failed,
Cancelled,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum ConnectPhase {
Dns,
TcpConnect,
TlsHandshake,
WsUpgrade,
}
#[derive(Debug, Clone, Serialize)]
#[non_exhaustive]
pub struct ConnectRecord {
pub schema_version: u32,
pub attempt_id: Uuid,
pub timestamp: String,
pub outcome: ConnectOutcome,
pub last_phase: ConnectPhase,
pub url: String,
pub connect_duration_ms: f64,
#[serde(skip_serializing_if = "Option::is_none")]
pub local_addr: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub peer_addr: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub dns_ms: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tcp_connect_ms: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tls_handshake_ms: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub ws_upgrade_ms: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub request_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub dg_error: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub error: Option<String>,
}
impl ConnectRecord {
pub fn sample() -> ConnectRecord {
ConnectRecord {
schema_version: SCHEMA_VERSION,
attempt_id: Uuid::nil(),
timestamp: "2026-01-01T00:00:00.000Z".to_string(),
outcome: ConnectOutcome::Completed,
last_phase: ConnectPhase::WsUpgrade,
url: "wss://api.deepgram.com/v1/listen".to_string(),
connect_duration_ms: 312.4,
local_addr: Some("10.0.4.17:53210".to_string()),
peer_addr: Some("203.0.113.5:443".to_string()),
dns_ms: Some(11.2),
tcp_connect_ms: Some(68.9),
tls_handshake_ms: Some(141.7),
ws_upgrade_ms: Some(90.6),
request_id: Some("00000000-0000-0000-0000-000000000000".to_string()),
dg_error: None,
error: None,
}
}
}
pub trait DiagnosticsSink: Send + Sync {
fn emit(&self, record: ConnectRecord);
}
impl DiagnosticsSink for tokio::sync::mpsc::UnboundedSender<ConnectRecord> {
fn emit(&self, record: ConnectRecord) {
let _ = self.send(record);
}
}
pub fn sink_fn<F>(f: F) -> impl DiagnosticsSink
where
F: Fn(ConnectRecord) + Send + Sync,
{
struct FnSink<F>(F);
impl<F> DiagnosticsSink for FnSink<F>
where
F: Fn(ConnectRecord) + Send + Sync,
{
fn emit(&self, record: ConnectRecord) {
(self.0)(record);
}
}
FnSink(f)
}
#[derive(Clone)]
pub(crate) struct SharedSink(pub(crate) Arc<dyn DiagnosticsSink>);
impl fmt::Debug for SharedSink {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("SharedSink(..)")
}
}
pub(crate) struct DiagnosticsGuard {
record: ConnectRecord,
start: Instant,
phase_start: Instant,
sink: SharedSink,
}
impl DiagnosticsGuard {
pub(crate) fn new(sink: SharedSink, url: &url::Url) -> Self {
let now = Instant::now();
DiagnosticsGuard {
record: ConnectRecord {
schema_version: SCHEMA_VERSION,
attempt_id: Uuid::new_v4(),
timestamp: rfc3339_utc(SystemTime::now()),
outcome: ConnectOutcome::Cancelled,
last_phase: ConnectPhase::Dns,
url: redacted_url(url),
connect_duration_ms: 0.0,
local_addr: None,
peer_addr: None,
dns_ms: None,
tcp_connect_ms: None,
tls_handshake_ms: None,
ws_upgrade_ms: None,
request_id: None,
dg_error: None,
error: None,
},
start: now,
phase_start: now,
sink,
}
}
fn enter_phase(&mut self, phase: ConnectPhase) {
self.record.last_phase = phase;
self.phase_start = Instant::now();
}
fn finish_phase(&mut self) {
let elapsed = ms(self.phase_start.elapsed());
let slot = match self.record.last_phase {
ConnectPhase::Dns => &mut self.record.dns_ms,
ConnectPhase::TcpConnect => &mut self.record.tcp_connect_ms,
ConnectPhase::TlsHandshake => &mut self.record.tls_handshake_ms,
ConnectPhase::WsUpgrade => &mut self.record.ws_upgrade_ms,
};
*slot = Some(elapsed);
}
fn set_addrs(&mut self, stream: &TcpStream) {
self.record.local_addr = stream.local_addr().ok().map(|a| a.to_string());
self.record.peer_addr = stream.peer_addr().ok().map(|a| a.to_string());
}
pub(crate) fn set_request_id(&mut self, request_id: &str) {
self.record.request_id = Some(request_id.to_string());
}
pub(crate) fn complete(&mut self) {
self.record.outcome = ConnectOutcome::Completed;
}
pub(crate) fn fail_client(&mut self, error: &str) {
self.record.outcome = ConnectOutcome::Failed;
self.record.error = Some(error.to_string());
}
fn fail(&mut self, error: &TungsteniteError) {
self.record.outcome = ConnectOutcome::Failed;
self.record.error = Some(error.to_string());
if let TungsteniteError::Http(response) = error {
if let Some(id) = header_str(response.headers(), "dg-request-id") {
self.record.request_id = Some(id);
}
if let Some(err) = header_str(response.headers(), "dg-error") {
self.record.dg_error = Some(err);
}
}
}
}
impl Drop for DiagnosticsGuard {
fn drop(&mut self) {
self.record.connect_duration_ms = ms(self.start.elapsed());
self.sink.0.emit(self.record.clone());
}
}
fn redacted_url(url: &url::Url) -> String {
let mut url = url.clone();
let _ = url.set_username("");
let _ = url.set_password(None);
url.set_query(None);
url.set_fragment(None);
url.to_string()
}
fn ms(duration: std::time::Duration) -> f64 {
duration.as_secs_f64() * 1000.0
}
fn header_str(headers: &http::HeaderMap, name: &str) -> Option<String> {
headers
.get(name)
.and_then(|v| v.to_str().ok())
.map(str::to_string)
}
pub(crate) async fn connect_with_diagnostics(
request: Request<()>,
guard: &mut DiagnosticsGuard,
) -> Result<(
WebSocketStream<MaybeTlsStream<TcpStream>>,
tungstenite::handshake::client::Response,
)> {
match connect_phases(request, guard).await {
Ok(ok) => Ok(ok),
Err(err) => {
guard.fail(&err);
Err(DeepgramError::from(Box::new(err)))
}
}
}
async fn connect_phases(
request: Request<()>,
guard: &mut DiagnosticsGuard,
) -> std::result::Result<
(
WebSocketStream<MaybeTlsStream<TcpStream>>,
tungstenite::handshake::client::Response,
),
TungsteniteError,
> {
let domain = domain(&request)?;
let tls = match request.uri().scheme_str() {
Some("wss") => true,
Some("ws") => false,
_ => return Err(TungsteniteError::Url(UrlError::UnsupportedUrlScheme)),
};
let port = request
.uri()
.port_u16()
.unwrap_or(if tls { 443 } else { 80 });
guard.enter_phase(ConnectPhase::Dns);
let addrs: Vec<_> = tokio::net::lookup_host((domain.as_str(), port))
.await?
.collect();
guard.finish_phase();
guard.enter_phase(ConnectPhase::TcpConnect);
let mut tcp = None;
let mut last_err = None;
for addr in addrs {
match TcpStream::connect(addr).await {
Ok(stream) => {
tcp = Some(stream);
break;
}
Err(err) => last_err = Some(err),
}
}
let tcp = match tcp {
Some(tcp) => tcp,
None => {
return Err(TungsteniteError::Io(last_err.unwrap_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::NotFound,
format!("could not resolve host {domain}"),
)
})))
}
};
guard.set_addrs(&tcp);
guard.finish_phase();
let stream = if tls {
guard.enter_phase(ConnectPhase::TlsHandshake);
let server_name = rustls_pki_types::ServerName::try_from(domain.as_str())
.map_err(|_| TungsteniteError::Tls(tungstenite::error::TlsError::InvalidDnsName))?
.to_owned();
let connector = tokio_rustls::TlsConnector::from(tls_client_config());
let tls_stream = connector
.connect(server_name, tcp)
.await
.map_err(TungsteniteError::Io)?;
guard.finish_phase();
MaybeTlsStream::Rustls(tls_stream)
} else {
MaybeTlsStream::Plain(tcp)
};
guard.enter_phase(ConnectPhase::WsUpgrade);
let (ws_stream, response) =
tokio_tungstenite::client_async_with_config(request, stream, None).await?;
guard.finish_phase();
Ok((ws_stream, response))
}
fn tls_client_config() -> Arc<rustls::ClientConfig> {
let mut root_store = rustls::RootCertStore::empty();
root_store.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
Arc::new(
rustls::ClientConfig::builder()
.with_root_certificates(root_store)
.with_no_client_auth(),
)
}
pub(crate) fn tls_connector() -> tokio_tungstenite::Connector {
tokio_tungstenite::Connector::Rustls(tls_client_config())
}
fn domain(request: &Request<()>) -> std::result::Result<String, TungsteniteError> {
match request.uri().host() {
Some(d) if d.starts_with('[') && d.ends_with(']') => Ok(d[1..d.len() - 1].to_string()),
Some(d) => Ok(d.to_string()),
None => Err(TungsteniteError::Url(UrlError::NoHostName)),
}
}
fn rfc3339_utc(time: SystemTime) -> String {
let duration = time.duration_since(UNIX_EPOCH).unwrap_or_default();
let secs = duration.as_secs();
let millis = duration.subsec_millis();
let days = (secs / 86_400) as i64;
let secs_of_day = secs % 86_400;
let (hour, minute, second) = (
secs_of_day / 3600,
(secs_of_day % 3600) / 60,
secs_of_day % 60,
);
let z = days + 719_468;
let era = z.div_euclid(146_097);
let doe = z.rem_euclid(146_097);
let yoe = (doe - doe / 1460 + doe / 36_524 - doe / 146_096) / 365;
let year = yoe + era * 400;
let doy = doe - (365 * yoe + yoe / 4 - yoe / 100);
let mp = (5 * doy + 2) / 153;
let day = doy - (153 * mp + 2) / 5 + 1;
let month = if mp < 10 { mp + 3 } else { mp - 9 };
let year = if month <= 2 { year + 1 } else { year };
format!("{year:04}-{month:02}-{day:02}T{hour:02}:{minute:02}:{second:02}.{millis:03}Z")
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
fn test_url() -> url::Url {
url::Url::parse("wss://api.deepgram.com/v1/listen?model=nova-3").unwrap()
}
fn channel_sink() -> (
SharedSink,
tokio::sync::mpsc::UnboundedReceiver<ConnectRecord>,
) {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
(SharedSink(Arc::new(tx)), rx)
}
#[test]
fn rfc3339_epoch() {
assert_eq!(rfc3339_utc(UNIX_EPOCH), "1970-01-01T00:00:00.000Z");
}
#[test]
fn rfc3339_known_instants() {
let t = UNIX_EPOCH + Duration::from_millis(1_789_049_002_114);
assert_eq!(rfc3339_utc(t), "2026-09-10T14:03:22.114Z");
let t = UNIX_EPOCH + Duration::from_millis(1_709_251_199_999);
assert_eq!(rfc3339_utc(t), "2024-02-29T23:59:59.999Z");
}
#[test]
fn record_serialization_skips_absent_fields() {
let (sink, mut rx) = channel_sink();
drop(DiagnosticsGuard::new(sink, &test_url()));
let record = rx.try_recv().expect("record emitted on drop");
let json: serde_json::Value =
serde_json::from_str(&serde_json::to_string(&record).unwrap()).unwrap();
assert_eq!(json["schema_version"], 1);
assert_eq!(json["outcome"], "cancelled");
assert_eq!(json["last_phase"], "dns");
assert_eq!(json["url"], "wss://api.deepgram.com/v1/listen");
for absent in [
"local_addr",
"peer_addr",
"dns_ms",
"tcp_connect_ms",
"tls_handshake_ms",
"ws_upgrade_ms",
"request_id",
"dg_error",
"error",
] {
assert!(
json.get(absent).is_none(),
"{absent} should be omitted when absent"
);
}
}
#[test]
fn recorded_url_omits_userinfo_and_query_values() {
let secret = "synthetic-secret-0f9b2";
let url = url::Url::parse(&format!(
"wss://user:{secret}@api.deepgram.com/v1/listen\
?model=nova-3&customer_token={secret}\
&callback=https%3A%2F%2Fexample.com%2Fcb%3Fsig%3D{secret}"
))
.unwrap();
let (sink, mut rx) = channel_sink();
drop(DiagnosticsGuard::new(sink, &url));
let record = rx.try_recv().expect("record emitted on drop");
assert_eq!(record.url, "wss://api.deepgram.com/v1/listen");
let line = serde_json::to_string(&record).unwrap();
assert!(
!line.contains(secret),
"serialized record must not contain query or userinfo secrets: {line}"
);
assert!(!line.contains("customer_token"));
assert!(!line.contains("model=nova-3"));
}
#[test]
fn cancelled_guard_keeps_finished_phase_timings() {
let (sink, mut rx) = channel_sink();
let mut guard = DiagnosticsGuard::new(sink, &test_url());
guard.enter_phase(ConnectPhase::Dns);
guard.finish_phase();
guard.enter_phase(ConnectPhase::TcpConnect);
guard.finish_phase();
guard.enter_phase(ConnectPhase::TlsHandshake);
drop(guard);
let record = rx.try_recv().expect("record emitted on drop");
assert_eq!(record.outcome, ConnectOutcome::Cancelled);
assert_eq!(record.last_phase, ConnectPhase::TlsHandshake);
assert!(record.dns_ms.is_some());
assert!(record.tcp_connect_ms.is_some());
assert!(record.tls_handshake_ms.is_none());
assert!(record.request_id.is_none());
}
#[tokio::test]
async fn record_survives_tokio_timeout() {
let (sink, mut rx) = channel_sink();
let connect = async {
let mut guard = DiagnosticsGuard::new(sink, &test_url());
guard.enter_phase(ConnectPhase::Dns);
guard.finish_phase();
guard.enter_phase(ConnectPhase::TcpConnect);
std::future::pending::<()>().await;
};
let result = tokio::time::timeout(Duration::from_millis(10), connect).await;
assert!(result.is_err(), "timeout should fire");
let record = rx.try_recv().expect("record emitted despite cancellation");
assert_eq!(record.outcome, ConnectOutcome::Cancelled);
assert_eq!(record.last_phase, ConnectPhase::TcpConnect);
assert!(record.connect_duration_ms >= 10.0);
}
#[test]
fn failed_upgrade_captures_deepgram_headers() {
let (sink, mut rx) = channel_sink();
let mut guard = DiagnosticsGuard::new(sink, &test_url());
guard.enter_phase(ConnectPhase::WsUpgrade);
let response = http::Response::builder()
.status(400)
.header("dg-request-id", "9ac2aaaa-bbbb-cccc-dddd-eeeeffff0000")
.header("dg-error", "bad model")
.body(None)
.unwrap();
guard.fail(&TungsteniteError::Http(Box::new(response)));
drop(guard);
let record = rx.try_recv().expect("record emitted on drop");
assert_eq!(record.outcome, ConnectOutcome::Failed);
assert_eq!(
record.request_id.as_deref(),
Some("9ac2aaaa-bbbb-cccc-dddd-eeeeffff0000")
);
assert_eq!(record.dg_error.as_deref(), Some("bad model"));
assert!(record.error.is_some());
}
#[test]
fn closure_sink_works() {
let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel();
let sink = SharedSink(Arc::new(sink_fn(move |record: ConnectRecord| {
let _ = tx.send(record.outcome);
})));
let mut guard = DiagnosticsGuard::new(sink, &test_url());
guard.complete();
drop(guard);
assert_eq!(rx.try_recv().unwrap(), ConnectOutcome::Completed);
}
#[test]
fn tls_config_matches_tokio_tungstenite_defaults() {
let config = tls_client_config();
assert!(!config.client_auth_cert_resolver.has_certs());
let roots = rustls::RootCertStore {
roots: webpki_roots::TLS_SERVER_ROOTS.to_vec(),
};
assert_eq!(
tls_client_config().crypto_provider().cipher_suites,
rustls::ClientConfig::builder()
.with_root_certificates(roots)
.with_no_client_auth()
.crypto_provider()
.cipher_suites,
);
}
}