1use rand::RngCore;
16use serde::{Deserialize, Serialize};
17use url::Url;
18
19use crate::client::Client;
20use crate::error::{Error, Result};
21use crate::request;
22
23#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
25pub enum CodeChallengeMethod {
26 #[serde(rename = "S256")]
28 S256,
29 #[serde(rename = "plain")]
31 Plain,
32}
33
34pub fn generate_code_verifier() -> String {
38 use base64::engine::general_purpose::URL_SAFE_NO_PAD;
39 use base64::Engine;
40 let mut buf = [0u8; 32];
41 rand::thread_rng().fill_bytes(&mut buf);
42 URL_SAFE_NO_PAD.encode(buf)
43}
44
45pub fn create_s256_code_challenge(verifier: &str) -> String {
48 use base64::engine::general_purpose::URL_SAFE_NO_PAD;
49 use base64::Engine;
50 use sha2::Digest;
51 let digest = sha2::Sha256::digest(verifier.as_bytes());
52 URL_SAFE_NO_PAD.encode(digest)
53}
54
55#[derive(Clone, Debug, Default, PartialEq, Eq)]
57pub struct AuthUrlParams<'a> {
58 pub callback_url: &'a str,
61 pub code_challenge: Option<&'a str>,
63 pub code_challenge_method: Option<CodeChallengeMethod>,
65}
66
67pub fn build_auth_url(base_url: &str, params: AuthUrlParams<'_>) -> Result<String> {
70 if params.callback_url.is_empty() {
71 return Err(Error::InvalidInput("callback_url is required"));
72 }
73 let mut u =
74 Url::parse(base_url).map_err(|_| Error::InvalidInput("base_url is not a valid URL"))?;
75 {
76 let mut q = u.query_pairs_mut();
77 q.append_pair("callback_url", params.callback_url);
78 if let Some(c) = params.code_challenge {
79 q.append_pair("code_challenge", c);
80 }
81 if let Some(m) = params.code_challenge_method {
82 let v = match m {
83 CodeChallengeMethod::S256 => "S256",
84 CodeChallengeMethod::Plain => "plain",
85 };
86 q.append_pair("code_challenge_method", v);
87 }
88 }
89 Ok(u.into())
90}
91
92#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
94pub struct ExchangeAuthCodeRequest {
95 pub code: String,
97 #[serde(skip_serializing_if = "Option::is_none", default)]
100 pub code_verifier: Option<String>,
101 #[serde(skip_serializing_if = "Option::is_none", default)]
103 pub code_challenge_method: Option<CodeChallengeMethod>,
104}
105
106#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
108pub struct ExchangeAuthCodeResponse {
109 #[serde(default)]
112 pub key: String,
113 #[serde(default)]
115 pub user_id: Option<String>,
116}
117
118impl Client {
119 pub async fn exchange_auth_code(
127 &self,
128 req: &ExchangeAuthCodeRequest,
129 ) -> Result<ExchangeAuthCodeResponse> {
130 if req.code.is_empty() {
131 return Err(Error::InvalidInput("code is required"));
132 }
133 request::execute_json(self, "auth/keys", req).await
134 }
135}
136
137#[cfg(test)]
138mod tests {
139 use super::*;
140
141 #[test]
142 fn verifier_is_43_chars_and_url_safe() {
143 let v = generate_code_verifier();
144 assert_eq!(v.len(), 43);
145 for c in v.chars() {
146 assert!(
147 c.is_ascii_alphanumeric() || c == '-' || c == '_',
148 "non-urlsafe char {c}"
149 );
150 }
151 }
152
153 #[test]
154 fn s256_challenge_known_vector() {
155 let verifier = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk";
157 let expected = "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM";
158 assert_eq!(create_s256_code_challenge(verifier), expected);
159 }
160
161 #[test]
162 fn build_auth_url_requires_callback() {
163 let err = build_auth_url(
164 "https://openrouter.ai/auth",
165 AuthUrlParams {
166 callback_url: "",
167 ..Default::default()
168 },
169 )
170 .unwrap_err();
171 assert!(matches!(err, Error::InvalidInput(_)));
172 }
173
174 #[test]
175 fn build_auth_url_appends_params() {
176 let url = build_auth_url(
177 "https://openrouter.ai/auth",
178 AuthUrlParams {
179 callback_url: "https://app.example/cb",
180 code_challenge: Some("CHAL"),
181 code_challenge_method: Some(CodeChallengeMethod::S256),
182 },
183 )
184 .unwrap();
185 assert!(url.contains("callback_url=https%3A%2F%2Fapp.example%2Fcb"));
186 assert!(url.contains("code_challenge=CHAL"));
187 assert!(url.contains("code_challenge_method=S256"));
188 }
189}