schwab_cli/
auth_callback.rs1use std::net::{IpAddr, Ipv4Addr};
2use std::sync::Arc;
3use std::time::Duration;
4
5use anyhow::{Context, Result};
6use rcgen::{
7 CertificateParams, DnType, ExtendedKeyUsagePurpose, KeyPair, KeyUsagePurpose, SanType,
8};
9use rustls::pki_types::{CertificateDer, PrivateKeyDer};
10use rustls::ServerConfig;
11use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
12use tokio::net::TcpListener;
13use tokio_rustls::TlsAcceptor;
14use tracing::debug;
15use url::Url;
16
17pub async fn capture_redirect_code(
22 redirect_uri: String,
23 timeout: Duration,
24) -> Result<Option<String>> {
25 let (use_tls, port) = parse_local_redirect(&redirect_uri)?;
26 let listener = TcpListener::bind(("127.0.0.1", port))
27 .await
28 .with_context(|| format!("Could not bind 127.0.0.1:{port} for OAuth callback"))?;
29
30 let acceptor = if use_tls {
31 crate::tls::install_crypto_provider();
32 Some(TlsAcceptor::from(build_tls_config()?))
33 } else {
34 None
35 };
36
37 let deadline = tokio::time::Instant::now() + timeout;
38
39 loop {
40 let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
41 if remaining.is_zero() {
42 return Ok(None);
43 }
44
45 let (stream, _) = match tokio::time::timeout(remaining, listener.accept()).await {
46 Ok(Ok(pair)) => pair,
47 Ok(Err(e)) => return Err(e.into()),
48 Err(_) => return Ok(None),
49 };
50
51 let code = if let Some(acceptor) = acceptor.clone() {
52 match acceptor.accept(stream).await {
53 Ok(tls) => handle_oauth_request(tls).await?,
54 Err(err) => {
55 debug!(%err, "TLS handshake failed; waiting for browser retry after cert acceptance");
56 continue;
57 }
58 }
59 } else {
60 handle_oauth_request(stream).await?
61 };
62
63 if code.is_some() {
64 return Ok(code);
65 }
66 }
67}
68
69async fn handle_oauth_request<S>(mut stream: S) -> Result<Option<String>>
70where
71 S: AsyncRead + AsyncWrite + Unpin,
72{
73 let mut buf = vec![0u8; 8192];
74 let n = stream.read(&mut buf).await?;
75 if n == 0 {
76 return Ok(None);
77 }
78
79 let request = String::from_utf8_lossy(&buf[..n]);
80 let path = request
81 .lines()
82 .next()
83 .and_then(|line| line.split_whitespace().nth(1))
84 .unwrap_or("/");
85
86 let code = Url::parse(&format!("https://127.0.0.1{path}"))
87 .ok()
88 .and_then(|u| {
89 u.query_pairs()
90 .find(|(k, _)| k == "code")
91 .map(|(_, v)| v.to_string())
92 });
93
94 let body = "Schwab OAuth complete. You can close this tab and return to the terminal.";
95 let response = format!(
96 "HTTP/1.1 200 OK\r\nContent-Type: text/plain; charset=utf-8\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
97 body.len()
98 );
99 let _ = stream.write_all(response.as_bytes()).await;
100 let _ = stream.shutdown().await;
101
102 Ok(code)
103}
104
105fn build_tls_config() -> Result<Arc<ServerConfig>> {
106 let key_pair = KeyPair::generate().context("Failed to generate TLS key pair")?;
107 let mut params = CertificateParams::default();
108 params
109 .distinguished_name
110 .push(DnType::CommonName, "127.0.0.1");
111 params.subject_alt_names = vec![
112 SanType::IpAddress(IpAddr::V4(Ipv4Addr::LOCALHOST)),
113 SanType::DnsName(
114 "localhost"
115 .try_into()
116 .context("Invalid localhost DNS SAN")?,
117 ),
118 ];
119 params.key_usages = vec![
120 KeyUsagePurpose::DigitalSignature,
121 KeyUsagePurpose::KeyEncipherment,
122 ];
123 params.extended_key_usages = vec![ExtendedKeyUsagePurpose::ServerAuth];
124
125 let cert = params
126 .self_signed(&key_pair)
127 .context("Failed to sign localhost certificate")?;
128 let cert_der = CertificateDer::from(cert.der().to_vec());
129 let key_der = PrivateKeyDer::Pkcs8(key_pair.serialize_der().into());
130
131 let config = ServerConfig::builder()
132 .with_no_client_auth()
133 .with_single_cert(vec![cert_der], key_der)
134 .context("Failed to build TLS server config")?;
135
136 Ok(Arc::new(config))
137}
138
139fn parse_local_redirect(redirect_uri: &str) -> Result<(bool, u16)> {
140 let url = Url::parse(redirect_uri).context("Invalid SCHWAB_REDIRECT_URI")?;
141 let host = url.host_str().unwrap_or("");
142 if host != "127.0.0.1" && host != "localhost" {
143 anyhow::bail!("Auto-capture only supports localhost redirect URIs");
144 }
145 let use_tls = url.scheme() == "https";
146 let port = url.port().unwrap_or(if use_tls { 443 } else { 80 });
147 Ok((use_tls, port))
148}
149
150#[cfg(test)]
151mod tests {
152 use super::*;
153
154 #[test]
155 fn parses_https_local_redirect() {
156 let (tls, port) = parse_local_redirect("https://127.0.0.1:8182").unwrap();
157 assert!(tls);
158 assert_eq!(port, 8182);
159 }
160
161 #[test]
162 fn parses_http_local_redirect() {
163 let (tls, port) = parse_local_redirect("http://127.0.0.1:8182/").unwrap();
164 assert!(!tls);
165 assert_eq!(port, 8182);
166 }
167}