Skip to content

Commit 6538e09

Browse files
committed
feat: enhance ClientAddr to include is_secure flag
1 parent ad7a4d7 commit 6538e09

2 files changed

Lines changed: 24 additions & 5 deletions

File tree

src/handler/middleware/clientaddr.rs

Lines changed: 19 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,17 +1,21 @@
11
use axum::extract::{ConnectInfo, FromRequestParts};
2-
use http::{request::Parts, StatusCode};
2+
use http::{StatusCode, request::Parts};
33
use std::{
44
fmt::{self, Formatter},
55
net::{IpAddr, Ipv4Addr, SocketAddr},
66
};
77

88
pub struct ClientAddr {
99
pub addr: SocketAddr,
10+
pub is_secure: bool,
1011
}
1112

1213
impl ClientAddr {
1314
pub fn new(addr: SocketAddr) -> Self {
14-
ClientAddr { addr }
15+
ClientAddr {
16+
addr,
17+
is_secure: false,
18+
}
1519
}
1620
pub fn ip(&self) -> IpAddr {
1721
self.addr.ip()
@@ -25,12 +29,20 @@ where
2529
type Rejection = StatusCode;
2630

2731
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
32+
let is_secure = match parts.uri.scheme_str() {
33+
Some("wss") | Some("https") => true,
34+
_ => parts
35+
.headers
36+
.get("x-forwarded-proto")
37+
.map_or(false, |v| v == "https"),
38+
};
2839
let mut remote_addr = match parts.extensions.get::<ConnectInfo<SocketAddr>>() {
2940
Some(ConnectInfo(addr)) => addr.clone(),
3041
None => {
3142
return Ok(ClientAddr {
3243
addr: SocketAddr::from((IpAddr::V4(Ipv4Addr::new(0, 0, 0, 0)), 0)),
33-
})
44+
is_secure,
45+
});
3446
}
3547
};
3648

@@ -49,7 +61,10 @@ where
4961
}
5062
}
5163
}
52-
Ok(ClientAddr { addr: remote_addr })
64+
Ok(ClientAddr {
65+
addr: remote_addr,
66+
is_secure,
67+
})
5368
}
5469
}
5570

src/proxy/ws.rs

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -20,7 +20,11 @@ pub async fn sip_ws_handler(
2020
let (from_ws_tx, from_ws_rx) = mpsc::unbounded_channel();
2121
let (to_ws_tx, mut to_ws_rx) = mpsc::unbounded_channel();
2222

23-
let transport_type = rsip::transport::Transport::Ws;
23+
let transport_type = if client_addr.is_secure {
24+
rsip::transport::Transport::Wss
25+
} else {
26+
rsip::transport::Transport::Ws
27+
};
2428
let local_addr = SipAddr {
2529
r#type: Some(transport_type),
2630
addr: client_addr.addr.clone().into(),

0 commit comments

Comments
 (0)