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