client.rs 39 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929930931932933934935936937938939940941942943944945946947948949950951952953954955956957958959960961962963964965966967968969970971972973974975976977978979980981982983984985986987988989990991992993994995996997998999100010011002100310041005100610071008100910101011101210131014101510161017101810191020102110221023102410251026102710281029103010311032103310341035103610371038103910401041104210431044104510461047104810491050105110521053105410551056105710581059106010611062106310641065106610671068106910701071107210731074107510761077107810791080108110821083108410851086108710881089109010911092109310941095109610971098109911001101110211031104110511061107110811091110111111121113111411151116111711181119112011211122112311241125112611271128112911301131113211331134113511361137113811391140114111421143114411451146114711481149115011511152115311541155115611571158115911601161116211631164116511661167116811691170117111721173
  1. use std::collections::VecDeque;
  2. use std::time::{Duration, SystemTime, UNIX_EPOCH};
  3. use runtime::{
  4. load_oauth_credentials, save_oauth_credentials, OAuthConfig, OAuthRefreshRequest,
  5. OAuthTokenExchangeRequest,
  6. };
  7. use serde::Deserialize;
  8. use serde_json::{Map, Value};
  9. use telemetry::{AnthropicRequestProfile, ClientIdentity, SessionTracer};
  10. use crate::error::ApiError;
  11. use crate::sse::SseParser;
  12. use crate::types::{MessageRequest, MessageResponse, StreamEvent};
  13. const DEFAULT_BASE_URL: &str = "https://api.anthropic.com";
  14. const MESSAGES_PATH: &str = "/v1/messages";
  15. const REQUEST_ID_HEADER: &str = "request-id";
  16. const ALT_REQUEST_ID_HEADER: &str = "x-request-id";
  17. const DEFAULT_INITIAL_BACKOFF: Duration = Duration::from_millis(200);
  18. const DEFAULT_MAX_BACKOFF: Duration = Duration::from_secs(2);
  19. const DEFAULT_MAX_RETRIES: u32 = 2;
  20. #[derive(Debug, Clone, PartialEq, Eq)]
  21. pub enum AuthSource {
  22. None,
  23. ApiKey(String),
  24. BearerToken(String),
  25. ApiKeyAndBearer {
  26. api_key: String,
  27. bearer_token: String,
  28. },
  29. }
  30. impl AuthSource {
  31. pub fn from_env() -> Result<Self, ApiError> {
  32. let api_key = read_env_non_empty("ANTHROPIC_API_KEY")?;
  33. let auth_token = read_env_non_empty("ANTHROPIC_AUTH_TOKEN")?;
  34. match (api_key, auth_token) {
  35. (Some(api_key), Some(bearer_token)) => Ok(Self::ApiKeyAndBearer {
  36. api_key,
  37. bearer_token,
  38. }),
  39. (Some(api_key), None) => Ok(Self::ApiKey(api_key)),
  40. (None, Some(bearer_token)) => Ok(Self::BearerToken(bearer_token)),
  41. (None, None) => Err(ApiError::MissingApiKey),
  42. }
  43. }
  44. #[must_use]
  45. pub fn api_key(&self) -> Option<&str> {
  46. match self {
  47. Self::ApiKey(api_key) | Self::ApiKeyAndBearer { api_key, .. } => Some(api_key),
  48. Self::None | Self::BearerToken(_) => None,
  49. }
  50. }
  51. #[must_use]
  52. pub fn bearer_token(&self) -> Option<&str> {
  53. match self {
  54. Self::BearerToken(token)
  55. | Self::ApiKeyAndBearer {
  56. bearer_token: token,
  57. ..
  58. } => Some(token),
  59. Self::None | Self::ApiKey(_) => None,
  60. }
  61. }
  62. #[must_use]
  63. pub fn masked_authorization_header(&self) -> &'static str {
  64. if self.bearer_token().is_some() {
  65. "Bearer [REDACTED]"
  66. } else {
  67. "<absent>"
  68. }
  69. }
  70. pub fn apply(&self, mut request_builder: reqwest::RequestBuilder) -> reqwest::RequestBuilder {
  71. if let Some(api_key) = self.api_key() {
  72. request_builder = request_builder.header("x-api-key", api_key);
  73. }
  74. if let Some(token) = self.bearer_token() {
  75. request_builder = request_builder.bearer_auth(token);
  76. }
  77. request_builder
  78. }
  79. }
  80. #[derive(Debug, Clone, PartialEq, Eq, Deserialize)]
  81. pub struct OAuthTokenSet {
  82. pub access_token: String,
  83. pub refresh_token: Option<String>,
  84. pub expires_at: Option<u64>,
  85. #[serde(default)]
  86. pub scopes: Vec<String>,
  87. }
  88. impl From<OAuthTokenSet> for AuthSource {
  89. fn from(value: OAuthTokenSet) -> Self {
  90. Self::BearerToken(value.access_token)
  91. }
  92. }
  93. #[derive(Debug, Clone)]
  94. pub struct AnthropicClient {
  95. http: reqwest::Client,
  96. auth: AuthSource,
  97. base_url: String,
  98. max_retries: u32,
  99. initial_backoff: Duration,
  100. max_backoff: Duration,
  101. request_profile: AnthropicRequestProfile,
  102. session_tracer: Option<SessionTracer>,
  103. }
  104. impl AnthropicClient {
  105. #[must_use]
  106. pub fn new(api_key: impl Into<String>) -> Self {
  107. Self {
  108. http: reqwest::Client::new(),
  109. auth: AuthSource::ApiKey(api_key.into()),
  110. base_url: DEFAULT_BASE_URL.to_string(),
  111. max_retries: DEFAULT_MAX_RETRIES,
  112. initial_backoff: DEFAULT_INITIAL_BACKOFF,
  113. max_backoff: DEFAULT_MAX_BACKOFF,
  114. request_profile: AnthropicRequestProfile::default(),
  115. session_tracer: None,
  116. }
  117. }
  118. #[must_use]
  119. pub fn from_auth(auth: AuthSource) -> Self {
  120. Self {
  121. http: reqwest::Client::new(),
  122. auth,
  123. base_url: DEFAULT_BASE_URL.to_string(),
  124. max_retries: DEFAULT_MAX_RETRIES,
  125. initial_backoff: DEFAULT_INITIAL_BACKOFF,
  126. max_backoff: DEFAULT_MAX_BACKOFF,
  127. request_profile: AnthropicRequestProfile::default(),
  128. session_tracer: None,
  129. }
  130. }
  131. pub fn from_env() -> Result<Self, ApiError> {
  132. Ok(Self::from_auth(AuthSource::from_env_or_saved()?).with_base_url(read_base_url()))
  133. }
  134. #[must_use]
  135. pub fn with_auth_source(mut self, auth: AuthSource) -> Self {
  136. self.auth = auth;
  137. self
  138. }
  139. #[must_use]
  140. pub fn with_auth_token(mut self, auth_token: Option<String>) -> Self {
  141. match (
  142. self.auth.api_key().map(ToOwned::to_owned),
  143. auth_token.filter(|token| !token.is_empty()),
  144. ) {
  145. (Some(api_key), Some(bearer_token)) => {
  146. self.auth = AuthSource::ApiKeyAndBearer {
  147. api_key,
  148. bearer_token,
  149. };
  150. }
  151. (Some(api_key), None) => {
  152. self.auth = AuthSource::ApiKey(api_key);
  153. }
  154. (None, Some(bearer_token)) => {
  155. self.auth = AuthSource::BearerToken(bearer_token);
  156. }
  157. (None, None) => {
  158. self.auth = AuthSource::None;
  159. }
  160. }
  161. self
  162. }
  163. #[must_use]
  164. pub fn with_base_url(mut self, base_url: impl Into<String>) -> Self {
  165. self.base_url = base_url.into();
  166. self
  167. }
  168. #[must_use]
  169. pub fn with_request_profile(mut self, request_profile: AnthropicRequestProfile) -> Self {
  170. self.request_profile = request_profile;
  171. self
  172. }
  173. #[must_use]
  174. pub fn with_client_identity(mut self, client_identity: ClientIdentity) -> Self {
  175. self.request_profile.client_identity = client_identity;
  176. self
  177. }
  178. #[must_use]
  179. pub fn with_beta(mut self, beta: impl Into<String>) -> Self {
  180. let beta = beta.into();
  181. if !self.request_profile.betas.contains(&beta) {
  182. self.request_profile.betas.push(beta);
  183. }
  184. self
  185. }
  186. #[must_use]
  187. pub fn with_extra_body_param(mut self, key: impl Into<String>, value: Value) -> Self {
  188. self.request_profile.extra_body.insert(key.into(), value);
  189. self
  190. }
  191. #[must_use]
  192. pub fn with_session_tracer(mut self, session_tracer: SessionTracer) -> Self {
  193. self.session_tracer = Some(session_tracer);
  194. self
  195. }
  196. #[must_use]
  197. pub fn with_retry_policy(
  198. mut self,
  199. max_retries: u32,
  200. initial_backoff: Duration,
  201. max_backoff: Duration,
  202. ) -> Self {
  203. self.max_retries = max_retries;
  204. self.initial_backoff = initial_backoff;
  205. self.max_backoff = max_backoff;
  206. self
  207. }
  208. #[must_use]
  209. pub fn auth_source(&self) -> &AuthSource {
  210. &self.auth
  211. }
  212. pub async fn send_message(
  213. &self,
  214. request: &MessageRequest,
  215. ) -> Result<MessageResponse, ApiError> {
  216. let request = MessageRequest {
  217. stream: false,
  218. ..request.clone()
  219. };
  220. let response = self.send_with_retry(&request).await?;
  221. let request_id = request_id_from_headers(response.headers());
  222. let mut response = response
  223. .json::<MessageResponse>()
  224. .await
  225. .map_err(ApiError::from)?;
  226. if response.request_id.is_none() {
  227. response.request_id = request_id;
  228. }
  229. Ok(response)
  230. }
  231. pub async fn stream_message(
  232. &self,
  233. request: &MessageRequest,
  234. ) -> Result<MessageStream, ApiError> {
  235. let response = self
  236. .send_with_retry(&request.clone().with_streaming())
  237. .await?;
  238. Ok(MessageStream {
  239. request_id: request_id_from_headers(response.headers()),
  240. response,
  241. parser: SseParser::new(),
  242. pending: VecDeque::new(),
  243. done: false,
  244. })
  245. }
  246. pub async fn exchange_oauth_code(
  247. &self,
  248. config: &OAuthConfig,
  249. request: &OAuthTokenExchangeRequest,
  250. ) -> Result<OAuthTokenSet, ApiError> {
  251. let response = self
  252. .http
  253. .post(&config.token_url)
  254. .header("content-type", "application/x-www-form-urlencoded")
  255. .form(&request.form_params())
  256. .send()
  257. .await
  258. .map_err(ApiError::from)?;
  259. let response = expect_success(response).await?;
  260. response
  261. .json::<OAuthTokenSet>()
  262. .await
  263. .map_err(ApiError::from)
  264. }
  265. pub async fn refresh_oauth_token(
  266. &self,
  267. config: &OAuthConfig,
  268. request: &OAuthRefreshRequest,
  269. ) -> Result<OAuthTokenSet, ApiError> {
  270. let response = self
  271. .http
  272. .post(&config.token_url)
  273. .header("content-type", "application/x-www-form-urlencoded")
  274. .form(&request.form_params())
  275. .send()
  276. .await
  277. .map_err(ApiError::from)?;
  278. let response = expect_success(response).await?;
  279. response
  280. .json::<OAuthTokenSet>()
  281. .await
  282. .map_err(ApiError::from)
  283. }
  284. async fn send_with_retry(
  285. &self,
  286. request: &MessageRequest,
  287. ) -> Result<reqwest::Response, ApiError> {
  288. let mut attempts = 0;
  289. let mut last_error: Option<ApiError>;
  290. loop {
  291. attempts += 1;
  292. self.record_request_started(request, attempts);
  293. match self.send_raw_request(request).await {
  294. Ok(response) => match expect_success(response).await {
  295. Ok(response) => {
  296. self.record_request_succeeded(request, attempts, &response);
  297. return Ok(response);
  298. }
  299. Err(error) if error.is_retryable() && attempts <= self.max_retries + 1 => {
  300. self.record_request_failed(request, attempts, &error);
  301. last_error = Some(error);
  302. }
  303. Err(error) => {
  304. self.record_request_failed(request, attempts, &error);
  305. return Err(error);
  306. }
  307. },
  308. Err(error) if error.is_retryable() && attempts <= self.max_retries + 1 => {
  309. self.record_request_failed(request, attempts, &error);
  310. last_error = Some(error);
  311. }
  312. Err(error) => {
  313. self.record_request_failed(request, attempts, &error);
  314. return Err(error);
  315. }
  316. }
  317. if attempts > self.max_retries {
  318. break;
  319. }
  320. tokio::time::sleep(self.backoff_for_attempt(attempts)?).await;
  321. }
  322. Err(ApiError::RetriesExhausted {
  323. attempts,
  324. last_error: Box::new(last_error.expect("retry loop must capture an error")),
  325. })
  326. }
  327. async fn send_raw_request(
  328. &self,
  329. request: &MessageRequest,
  330. ) -> Result<reqwest::Response, ApiError> {
  331. let request_url = format!("{}{}", self.base_url.trim_end_matches('/'), MESSAGES_PATH);
  332. let mut request_builder = self
  333. .http
  334. .post(&request_url)
  335. .header("content-type", "application/json");
  336. for (name, value) in self.request_profile.header_pairs() {
  337. request_builder = request_builder.header(name, value);
  338. }
  339. let mut request_builder = self.auth.apply(request_builder);
  340. let request_body = self.request_profile.render_json_body(request)?;
  341. request_builder = request_builder.json(&request_body);
  342. request_builder.send().await.map_err(ApiError::from)
  343. }
  344. fn record_request_started(&self, request: &MessageRequest, attempt: u32) {
  345. if let Some(tracer) = &self.session_tracer {
  346. tracer.record_http_request_started(
  347. attempt,
  348. "POST",
  349. MESSAGES_PATH,
  350. self.request_attributes(request),
  351. );
  352. }
  353. }
  354. fn record_request_succeeded(
  355. &self,
  356. request: &MessageRequest,
  357. attempt: u32,
  358. response: &reqwest::Response,
  359. ) {
  360. if let Some(tracer) = &self.session_tracer {
  361. tracer.record_http_request_succeeded(
  362. attempt,
  363. "POST",
  364. MESSAGES_PATH,
  365. response.status().as_u16(),
  366. request_id_from_headers(response.headers()),
  367. self.request_attributes(request),
  368. );
  369. }
  370. }
  371. fn record_request_failed(&self, request: &MessageRequest, attempt: u32, error: &ApiError) {
  372. if let Some(tracer) = &self.session_tracer {
  373. tracer.record_http_request_failed(
  374. attempt,
  375. "POST",
  376. MESSAGES_PATH,
  377. error.to_string(),
  378. error.is_retryable(),
  379. self.error_attributes(request, error),
  380. );
  381. }
  382. }
  383. fn request_attributes(&self, request: &MessageRequest) -> Map<String, Value> {
  384. let mut attributes = Map::new();
  385. attributes.insert("model".to_string(), Value::String(request.model.clone()));
  386. attributes.insert("stream".to_string(), Value::Bool(request.stream));
  387. attributes.insert("max_tokens".to_string(), Value::from(request.max_tokens));
  388. attributes.insert(
  389. "message_count".to_string(),
  390. Value::from(u64::try_from(request.messages.len()).unwrap_or(u64::MAX)),
  391. );
  392. attributes.insert(
  393. "tool_count".to_string(),
  394. Value::from(
  395. u64::try_from(request.tools.as_ref().map_or(0, Vec::len)).unwrap_or(u64::MAX),
  396. ),
  397. );
  398. attributes.insert(
  399. "beta_count".to_string(),
  400. Value::from(u64::try_from(self.request_profile.betas.len()).unwrap_or(u64::MAX)),
  401. );
  402. if !self.request_profile.betas.is_empty() {
  403. attributes.insert(
  404. "betas".to_string(),
  405. Value::Array(
  406. self.request_profile
  407. .betas
  408. .iter()
  409. .cloned()
  410. .map(Value::String)
  411. .collect(),
  412. ),
  413. );
  414. }
  415. if !self.request_profile.extra_body.is_empty() {
  416. attributes.insert(
  417. "extra_body_keys".to_string(),
  418. Value::Array(
  419. self.request_profile
  420. .extra_body
  421. .keys()
  422. .cloned()
  423. .map(Value::String)
  424. .collect(),
  425. ),
  426. );
  427. }
  428. attributes
  429. }
  430. fn error_attributes(&self, request: &MessageRequest, error: &ApiError) -> Map<String, Value> {
  431. let mut attributes = self.request_attributes(request);
  432. match error {
  433. ApiError::Api {
  434. status,
  435. error_type,
  436. message,
  437. ..
  438. } => {
  439. attributes.insert("status".to_string(), Value::from(status.as_u16()));
  440. if let Some(error_type) = error_type {
  441. attributes.insert("error_type".to_string(), Value::String(error_type.clone()));
  442. }
  443. if let Some(message) = message {
  444. attributes.insert("api_message".to_string(), Value::String(message.clone()));
  445. }
  446. }
  447. ApiError::Http(_) => {
  448. attributes.insert("error_type".to_string(), Value::String("http".to_string()));
  449. }
  450. ApiError::Json(_) => {
  451. attributes.insert("error_type".to_string(), Value::String("json".to_string()));
  452. }
  453. _ => {
  454. attributes.insert(
  455. "error_type".to_string(),
  456. Value::String("client".to_string()),
  457. );
  458. }
  459. }
  460. attributes
  461. }
  462. fn backoff_for_attempt(&self, attempt: u32) -> Result<Duration, ApiError> {
  463. let Some(multiplier) = 1_u32.checked_shl(attempt.saturating_sub(1)) else {
  464. return Err(ApiError::BackoffOverflow {
  465. attempt,
  466. base_delay: self.initial_backoff,
  467. });
  468. };
  469. Ok(self
  470. .initial_backoff
  471. .checked_mul(multiplier)
  472. .map_or(self.max_backoff, |delay| delay.min(self.max_backoff)))
  473. }
  474. }
  475. impl AuthSource {
  476. pub fn from_env_or_saved() -> Result<Self, ApiError> {
  477. if let Some(api_key) = read_env_non_empty("ANTHROPIC_API_KEY")? {
  478. return match read_env_non_empty("ANTHROPIC_AUTH_TOKEN")? {
  479. Some(bearer_token) => Ok(Self::ApiKeyAndBearer {
  480. api_key,
  481. bearer_token,
  482. }),
  483. None => Ok(Self::ApiKey(api_key)),
  484. };
  485. }
  486. if let Some(bearer_token) = read_env_non_empty("ANTHROPIC_AUTH_TOKEN")? {
  487. return Ok(Self::BearerToken(bearer_token));
  488. }
  489. match load_saved_oauth_token() {
  490. Ok(Some(token_set)) if oauth_token_is_expired(&token_set) => {
  491. if token_set.refresh_token.is_some() {
  492. Err(ApiError::Auth(
  493. "saved OAuth token is expired; load runtime OAuth config to refresh it"
  494. .to_string(),
  495. ))
  496. } else {
  497. Err(ApiError::ExpiredOAuthToken)
  498. }
  499. }
  500. Ok(Some(token_set)) => Ok(Self::BearerToken(token_set.access_token)),
  501. Ok(None) => Err(ApiError::MissingApiKey),
  502. Err(error) => Err(error),
  503. }
  504. }
  505. }
  506. #[must_use]
  507. pub fn oauth_token_is_expired(token_set: &OAuthTokenSet) -> bool {
  508. token_set
  509. .expires_at
  510. .is_some_and(|expires_at| expires_at <= now_unix_timestamp())
  511. }
  512. pub fn resolve_saved_oauth_token(config: &OAuthConfig) -> Result<Option<OAuthTokenSet>, ApiError> {
  513. let Some(token_set) = load_saved_oauth_token()? else {
  514. return Ok(None);
  515. };
  516. resolve_saved_oauth_token_set(config, token_set).map(Some)
  517. }
  518. pub fn resolve_startup_auth_source<F>(load_oauth_config: F) -> Result<AuthSource, ApiError>
  519. where
  520. F: FnOnce() -> Result<Option<OAuthConfig>, ApiError>,
  521. {
  522. if let Some(api_key) = read_env_non_empty("ANTHROPIC_API_KEY")? {
  523. return match read_env_non_empty("ANTHROPIC_AUTH_TOKEN")? {
  524. Some(bearer_token) => Ok(AuthSource::ApiKeyAndBearer {
  525. api_key,
  526. bearer_token,
  527. }),
  528. None => Ok(AuthSource::ApiKey(api_key)),
  529. };
  530. }
  531. if let Some(bearer_token) = read_env_non_empty("ANTHROPIC_AUTH_TOKEN")? {
  532. return Ok(AuthSource::BearerToken(bearer_token));
  533. }
  534. let Some(token_set) = load_saved_oauth_token()? else {
  535. return Err(ApiError::MissingApiKey);
  536. };
  537. if !oauth_token_is_expired(&token_set) {
  538. return Ok(AuthSource::BearerToken(token_set.access_token));
  539. }
  540. if token_set.refresh_token.is_none() {
  541. return Err(ApiError::ExpiredOAuthToken);
  542. }
  543. let Some(config) = load_oauth_config()? else {
  544. return Err(ApiError::Auth(
  545. "saved OAuth token is expired; runtime OAuth config is missing".to_string(),
  546. ));
  547. };
  548. Ok(AuthSource::from(resolve_saved_oauth_token_set(
  549. &config, token_set,
  550. )?))
  551. }
  552. fn resolve_saved_oauth_token_set(
  553. config: &OAuthConfig,
  554. token_set: OAuthTokenSet,
  555. ) -> Result<OAuthTokenSet, ApiError> {
  556. if !oauth_token_is_expired(&token_set) {
  557. return Ok(token_set);
  558. }
  559. let Some(refresh_token) = token_set.refresh_token.clone() else {
  560. return Err(ApiError::ExpiredOAuthToken);
  561. };
  562. let client = AnthropicClient::from_auth(AuthSource::None).with_base_url(read_base_url());
  563. let refreshed = client_runtime_block_on(async {
  564. client
  565. .refresh_oauth_token(
  566. config,
  567. &OAuthRefreshRequest::from_config(
  568. config,
  569. refresh_token,
  570. Some(token_set.scopes.clone()),
  571. ),
  572. )
  573. .await
  574. })?;
  575. let resolved = OAuthTokenSet {
  576. access_token: refreshed.access_token,
  577. refresh_token: refreshed.refresh_token.or(token_set.refresh_token),
  578. expires_at: refreshed.expires_at,
  579. scopes: refreshed.scopes,
  580. };
  581. save_oauth_credentials(&runtime::OAuthTokenSet {
  582. access_token: resolved.access_token.clone(),
  583. refresh_token: resolved.refresh_token.clone(),
  584. expires_at: resolved.expires_at,
  585. scopes: resolved.scopes.clone(),
  586. })
  587. .map_err(ApiError::from)?;
  588. Ok(resolved)
  589. }
  590. fn client_runtime_block_on<F, T>(future: F) -> Result<T, ApiError>
  591. where
  592. F: std::future::Future<Output = Result<T, ApiError>>,
  593. {
  594. tokio::runtime::Runtime::new()
  595. .map_err(ApiError::from)?
  596. .block_on(future)
  597. }
  598. fn load_saved_oauth_token() -> Result<Option<OAuthTokenSet>, ApiError> {
  599. let token_set = load_oauth_credentials().map_err(ApiError::from)?;
  600. Ok(token_set.map(|token_set| OAuthTokenSet {
  601. access_token: token_set.access_token,
  602. refresh_token: token_set.refresh_token,
  603. expires_at: token_set.expires_at,
  604. scopes: token_set.scopes,
  605. }))
  606. }
  607. fn now_unix_timestamp() -> u64 {
  608. SystemTime::now()
  609. .duration_since(UNIX_EPOCH)
  610. .map_or(0, |duration| duration.as_secs())
  611. }
  612. fn read_env_non_empty(key: &str) -> Result<Option<String>, ApiError> {
  613. match std::env::var(key) {
  614. Ok(value) if !value.is_empty() => Ok(Some(value)),
  615. Ok(_) | Err(std::env::VarError::NotPresent) => Ok(None),
  616. Err(error) => Err(ApiError::from(error)),
  617. }
  618. }
  619. #[cfg(test)]
  620. fn read_api_key() -> Result<String, ApiError> {
  621. let auth = AuthSource::from_env_or_saved()?;
  622. auth.api_key()
  623. .or_else(|| auth.bearer_token())
  624. .map(ToOwned::to_owned)
  625. .ok_or(ApiError::MissingApiKey)
  626. }
  627. #[cfg(test)]
  628. fn read_auth_token() -> Option<String> {
  629. read_env_non_empty("ANTHROPIC_AUTH_TOKEN")
  630. .ok()
  631. .and_then(std::convert::identity)
  632. }
  633. #[must_use]
  634. pub fn read_base_url() -> String {
  635. std::env::var("ANTHROPIC_BASE_URL").unwrap_or_else(|_| DEFAULT_BASE_URL.to_string())
  636. }
  637. fn request_id_from_headers(headers: &reqwest::header::HeaderMap) -> Option<String> {
  638. headers
  639. .get(REQUEST_ID_HEADER)
  640. .or_else(|| headers.get(ALT_REQUEST_ID_HEADER))
  641. .and_then(|value| value.to_str().ok())
  642. .map(ToOwned::to_owned)
  643. }
  644. #[derive(Debug)]
  645. pub struct MessageStream {
  646. request_id: Option<String>,
  647. response: reqwest::Response,
  648. parser: SseParser,
  649. pending: VecDeque<StreamEvent>,
  650. done: bool,
  651. }
  652. impl MessageStream {
  653. #[must_use]
  654. pub fn request_id(&self) -> Option<&str> {
  655. self.request_id.as_deref()
  656. }
  657. pub async fn next_event(&mut self) -> Result<Option<StreamEvent>, ApiError> {
  658. loop {
  659. if let Some(event) = self.pending.pop_front() {
  660. return Ok(Some(event));
  661. }
  662. if self.done {
  663. let remaining = self.parser.finish()?;
  664. self.pending.extend(remaining);
  665. if let Some(event) = self.pending.pop_front() {
  666. return Ok(Some(event));
  667. }
  668. return Ok(None);
  669. }
  670. match self.response.chunk().await? {
  671. Some(chunk) => {
  672. self.pending.extend(self.parser.push(&chunk)?);
  673. }
  674. None => {
  675. self.done = true;
  676. }
  677. }
  678. }
  679. }
  680. }
  681. async fn expect_success(response: reqwest::Response) -> Result<reqwest::Response, ApiError> {
  682. let status = response.status();
  683. if status.is_success() {
  684. return Ok(response);
  685. }
  686. let body = response.text().await.unwrap_or_else(|_| String::new());
  687. let parsed_error = serde_json::from_str::<AnthropicErrorEnvelope>(&body).ok();
  688. let retryable = is_retryable_status(status);
  689. Err(ApiError::Api {
  690. status,
  691. error_type: parsed_error
  692. .as_ref()
  693. .map(|error| error.error.error_type.clone()),
  694. message: parsed_error
  695. .as_ref()
  696. .map(|error| error.error.message.clone()),
  697. body,
  698. retryable,
  699. })
  700. }
  701. const fn is_retryable_status(status: reqwest::StatusCode) -> bool {
  702. matches!(status.as_u16(), 408 | 409 | 429 | 500 | 502 | 503 | 504)
  703. }
  704. #[derive(Debug, Deserialize)]
  705. struct AnthropicErrorEnvelope {
  706. error: AnthropicErrorBody,
  707. }
  708. #[derive(Debug, Deserialize)]
  709. struct AnthropicErrorBody {
  710. #[serde(rename = "type")]
  711. error_type: String,
  712. message: String,
  713. }
  714. #[cfg(test)]
  715. mod tests {
  716. use super::{ALT_REQUEST_ID_HEADER, REQUEST_ID_HEADER};
  717. use std::io::{Read, Write};
  718. use std::net::TcpListener;
  719. use std::sync::{Mutex, OnceLock};
  720. use std::thread;
  721. use std::time::{Duration, SystemTime, UNIX_EPOCH};
  722. use runtime::{clear_oauth_credentials, save_oauth_credentials, OAuthConfig};
  723. use crate::client::{
  724. now_unix_timestamp, oauth_token_is_expired, resolve_saved_oauth_token,
  725. resolve_startup_auth_source, AnthropicClient, AuthSource, OAuthTokenSet,
  726. };
  727. use crate::types::{ContentBlockDelta, MessageRequest};
  728. fn env_lock() -> std::sync::MutexGuard<'static, ()> {
  729. static LOCK: OnceLock<Mutex<()>> = OnceLock::new();
  730. LOCK.get_or_init(|| Mutex::new(()))
  731. .lock()
  732. .expect("env lock")
  733. }
  734. fn temp_config_home() -> std::path::PathBuf {
  735. std::env::temp_dir().join(format!(
  736. "api-oauth-test-{}-{}",
  737. std::process::id(),
  738. SystemTime::now()
  739. .duration_since(UNIX_EPOCH)
  740. .expect("time")
  741. .as_nanos()
  742. ))
  743. }
  744. fn sample_oauth_config(token_url: String) -> OAuthConfig {
  745. OAuthConfig {
  746. client_id: "runtime-client".to_string(),
  747. authorize_url: "https://console.test/oauth/authorize".to_string(),
  748. token_url,
  749. callback_port: Some(4545),
  750. manual_redirect_url: Some("https://console.test/oauth/callback".to_string()),
  751. scopes: vec!["org:read".to_string(), "user:write".to_string()],
  752. }
  753. }
  754. fn spawn_token_server(response_body: &'static str) -> String {
  755. let listener = TcpListener::bind("127.0.0.1:0").expect("bind listener");
  756. let address = listener.local_addr().expect("local addr");
  757. thread::spawn(move || {
  758. let (mut stream, _) = listener.accept().expect("accept connection");
  759. let mut buffer = [0_u8; 4096];
  760. let _ = stream.read(&mut buffer).expect("read request");
  761. let response = format!(
  762. "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\n\r\n{}",
  763. response_body.len(),
  764. response_body
  765. );
  766. stream
  767. .write_all(response.as_bytes())
  768. .expect("write response");
  769. });
  770. format!("http://{address}/oauth/token")
  771. }
  772. #[test]
  773. fn read_api_key_requires_presence() {
  774. let _guard = env_lock();
  775. std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
  776. std::env::remove_var("ANTHROPIC_API_KEY");
  777. std::env::remove_var("CLAUDE_CONFIG_HOME");
  778. let error = super::read_api_key().expect_err("missing key should error");
  779. assert!(matches!(error, crate::error::ApiError::MissingApiKey));
  780. }
  781. #[test]
  782. fn read_api_key_requires_non_empty_value() {
  783. let _guard = env_lock();
  784. std::env::set_var("ANTHROPIC_AUTH_TOKEN", "");
  785. std::env::remove_var("ANTHROPIC_API_KEY");
  786. let error = super::read_api_key().expect_err("empty key should error");
  787. assert!(matches!(error, crate::error::ApiError::MissingApiKey));
  788. std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
  789. }
  790. #[test]
  791. fn read_api_key_prefers_api_key_env() {
  792. let _guard = env_lock();
  793. std::env::set_var("ANTHROPIC_AUTH_TOKEN", "auth-token");
  794. std::env::set_var("ANTHROPIC_API_KEY", "legacy-key");
  795. assert_eq!(
  796. super::read_api_key().expect("api key should load"),
  797. "legacy-key"
  798. );
  799. std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
  800. std::env::remove_var("ANTHROPIC_API_KEY");
  801. }
  802. #[test]
  803. fn read_auth_token_reads_auth_token_env() {
  804. let _guard = env_lock();
  805. std::env::set_var("ANTHROPIC_AUTH_TOKEN", "auth-token");
  806. assert_eq!(super::read_auth_token().as_deref(), Some("auth-token"));
  807. std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
  808. }
  809. #[test]
  810. fn oauth_token_maps_to_bearer_auth_source() {
  811. let auth = AuthSource::from(OAuthTokenSet {
  812. access_token: "access-token".to_string(),
  813. refresh_token: Some("refresh".to_string()),
  814. expires_at: Some(123),
  815. scopes: vec!["scope:a".to_string()],
  816. });
  817. assert_eq!(auth.bearer_token(), Some("access-token"));
  818. assert_eq!(auth.api_key(), None);
  819. }
  820. #[test]
  821. fn auth_source_from_env_combines_api_key_and_bearer_token() {
  822. let _guard = env_lock();
  823. std::env::set_var("ANTHROPIC_AUTH_TOKEN", "auth-token");
  824. std::env::set_var("ANTHROPIC_API_KEY", "legacy-key");
  825. let auth = AuthSource::from_env().expect("env auth");
  826. assert_eq!(auth.api_key(), Some("legacy-key"));
  827. assert_eq!(auth.bearer_token(), Some("auth-token"));
  828. std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
  829. std::env::remove_var("ANTHROPIC_API_KEY");
  830. }
  831. #[test]
  832. fn auth_source_from_saved_oauth_when_env_absent() {
  833. let _guard = env_lock();
  834. let config_home = temp_config_home();
  835. std::env::set_var("CLAUDE_CONFIG_HOME", &config_home);
  836. std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
  837. std::env::remove_var("ANTHROPIC_API_KEY");
  838. save_oauth_credentials(&runtime::OAuthTokenSet {
  839. access_token: "saved-access-token".to_string(),
  840. refresh_token: Some("refresh".to_string()),
  841. expires_at: Some(now_unix_timestamp() + 300),
  842. scopes: vec!["scope:a".to_string()],
  843. })
  844. .expect("save oauth credentials");
  845. let auth = AuthSource::from_env_or_saved().expect("saved auth");
  846. assert_eq!(auth.bearer_token(), Some("saved-access-token"));
  847. clear_oauth_credentials().expect("clear credentials");
  848. std::env::remove_var("CLAUDE_CONFIG_HOME");
  849. std::fs::remove_dir_all(config_home).expect("cleanup temp dir");
  850. }
  851. #[test]
  852. fn oauth_token_expiry_uses_expires_at_timestamp() {
  853. assert!(oauth_token_is_expired(&OAuthTokenSet {
  854. access_token: "access-token".to_string(),
  855. refresh_token: None,
  856. expires_at: Some(1),
  857. scopes: Vec::new(),
  858. }));
  859. assert!(!oauth_token_is_expired(&OAuthTokenSet {
  860. access_token: "access-token".to_string(),
  861. refresh_token: None,
  862. expires_at: Some(now_unix_timestamp() + 60),
  863. scopes: Vec::new(),
  864. }));
  865. }
  866. #[test]
  867. fn resolve_saved_oauth_token_refreshes_expired_credentials() {
  868. let _guard = env_lock();
  869. let config_home = temp_config_home();
  870. std::env::set_var("CLAUDE_CONFIG_HOME", &config_home);
  871. std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
  872. std::env::remove_var("ANTHROPIC_API_KEY");
  873. save_oauth_credentials(&runtime::OAuthTokenSet {
  874. access_token: "expired-access-token".to_string(),
  875. refresh_token: Some("refresh-token".to_string()),
  876. expires_at: Some(1),
  877. scopes: vec!["scope:a".to_string()],
  878. })
  879. .expect("save expired oauth credentials");
  880. let token_url = spawn_token_server(
  881. "{\"access_token\":\"refreshed-token\",\"refresh_token\":\"fresh-refresh\",\"expires_at\":9999999999,\"scopes\":[\"scope:a\"]}",
  882. );
  883. let resolved = resolve_saved_oauth_token(&sample_oauth_config(token_url))
  884. .expect("resolve refreshed token")
  885. .expect("token set present");
  886. assert_eq!(resolved.access_token, "refreshed-token");
  887. let stored = runtime::load_oauth_credentials()
  888. .expect("load stored credentials")
  889. .expect("stored token set");
  890. assert_eq!(stored.access_token, "refreshed-token");
  891. clear_oauth_credentials().expect("clear credentials");
  892. std::env::remove_var("CLAUDE_CONFIG_HOME");
  893. std::fs::remove_dir_all(config_home).expect("cleanup temp dir");
  894. }
  895. #[test]
  896. fn resolve_startup_auth_source_uses_saved_oauth_without_loading_config() {
  897. let _guard = env_lock();
  898. let config_home = temp_config_home();
  899. std::env::set_var("CLAUDE_CONFIG_HOME", &config_home);
  900. std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
  901. std::env::remove_var("ANTHROPIC_API_KEY");
  902. save_oauth_credentials(&runtime::OAuthTokenSet {
  903. access_token: "saved-access-token".to_string(),
  904. refresh_token: Some("refresh".to_string()),
  905. expires_at: Some(now_unix_timestamp() + 300),
  906. scopes: vec!["scope:a".to_string()],
  907. })
  908. .expect("save oauth credentials");
  909. let auth = resolve_startup_auth_source(|| panic!("config should not be loaded"))
  910. .expect("startup auth");
  911. assert_eq!(auth.bearer_token(), Some("saved-access-token"));
  912. clear_oauth_credentials().expect("clear credentials");
  913. std::env::remove_var("CLAUDE_CONFIG_HOME");
  914. std::fs::remove_dir_all(config_home).expect("cleanup temp dir");
  915. }
  916. #[test]
  917. fn resolve_startup_auth_source_errors_when_refreshable_token_lacks_config() {
  918. let _guard = env_lock();
  919. let config_home = temp_config_home();
  920. std::env::set_var("CLAUDE_CONFIG_HOME", &config_home);
  921. std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
  922. std::env::remove_var("ANTHROPIC_API_KEY");
  923. save_oauth_credentials(&runtime::OAuthTokenSet {
  924. access_token: "expired-access-token".to_string(),
  925. refresh_token: Some("refresh-token".to_string()),
  926. expires_at: Some(1),
  927. scopes: vec!["scope:a".to_string()],
  928. })
  929. .expect("save expired oauth credentials");
  930. let error =
  931. resolve_startup_auth_source(|| Ok(None)).expect_err("missing config should error");
  932. assert!(
  933. matches!(error, crate::error::ApiError::Auth(message) if message.contains("runtime OAuth config is missing"))
  934. );
  935. let stored = runtime::load_oauth_credentials()
  936. .expect("load stored credentials")
  937. .expect("stored token set");
  938. assert_eq!(stored.access_token, "expired-access-token");
  939. assert_eq!(stored.refresh_token.as_deref(), Some("refresh-token"));
  940. clear_oauth_credentials().expect("clear credentials");
  941. std::env::remove_var("CLAUDE_CONFIG_HOME");
  942. std::fs::remove_dir_all(config_home).expect("cleanup temp dir");
  943. }
  944. #[test]
  945. fn resolve_saved_oauth_token_preserves_refresh_token_when_refresh_response_omits_it() {
  946. let _guard = env_lock();
  947. let config_home = temp_config_home();
  948. std::env::set_var("CLAUDE_CONFIG_HOME", &config_home);
  949. std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
  950. std::env::remove_var("ANTHROPIC_API_KEY");
  951. save_oauth_credentials(&runtime::OAuthTokenSet {
  952. access_token: "expired-access-token".to_string(),
  953. refresh_token: Some("refresh-token".to_string()),
  954. expires_at: Some(1),
  955. scopes: vec!["scope:a".to_string()],
  956. })
  957. .expect("save expired oauth credentials");
  958. let token_url = spawn_token_server(
  959. "{\"access_token\":\"refreshed-token\",\"expires_at\":9999999999,\"scopes\":[\"scope:a\"]}",
  960. );
  961. let resolved = resolve_saved_oauth_token(&sample_oauth_config(token_url))
  962. .expect("resolve refreshed token")
  963. .expect("token set present");
  964. assert_eq!(resolved.access_token, "refreshed-token");
  965. assert_eq!(resolved.refresh_token.as_deref(), Some("refresh-token"));
  966. let stored = runtime::load_oauth_credentials()
  967. .expect("load stored credentials")
  968. .expect("stored token set");
  969. assert_eq!(stored.refresh_token.as_deref(), Some("refresh-token"));
  970. clear_oauth_credentials().expect("clear credentials");
  971. std::env::remove_var("CLAUDE_CONFIG_HOME");
  972. std::fs::remove_dir_all(config_home).expect("cleanup temp dir");
  973. }
  974. #[test]
  975. fn message_request_stream_helper_sets_stream_true() {
  976. let request = MessageRequest {
  977. model: "claude-opus-4-6".to_string(),
  978. max_tokens: 64,
  979. messages: vec![],
  980. system: None,
  981. tools: None,
  982. tool_choice: None,
  983. stream: false,
  984. };
  985. assert!(request.with_streaming().stream);
  986. }
  987. #[test]
  988. fn backoff_doubles_until_maximum() {
  989. let client = AnthropicClient::new("test-key").with_retry_policy(
  990. 3,
  991. Duration::from_millis(10),
  992. Duration::from_millis(25),
  993. );
  994. assert_eq!(
  995. client.backoff_for_attempt(1).expect("attempt 1"),
  996. Duration::from_millis(10)
  997. );
  998. assert_eq!(
  999. client.backoff_for_attempt(2).expect("attempt 2"),
  1000. Duration::from_millis(20)
  1001. );
  1002. assert_eq!(
  1003. client.backoff_for_attempt(3).expect("attempt 3"),
  1004. Duration::from_millis(25)
  1005. );
  1006. }
  1007. #[test]
  1008. fn retryable_statuses_are_detected() {
  1009. assert!(super::is_retryable_status(
  1010. reqwest::StatusCode::TOO_MANY_REQUESTS
  1011. ));
  1012. assert!(super::is_retryable_status(
  1013. reqwest::StatusCode::INTERNAL_SERVER_ERROR
  1014. ));
  1015. assert!(!super::is_retryable_status(
  1016. reqwest::StatusCode::UNAUTHORIZED
  1017. ));
  1018. }
  1019. #[test]
  1020. fn tool_delta_variant_round_trips() {
  1021. let delta = ContentBlockDelta::InputJsonDelta {
  1022. partial_json: "{\"city\":\"Paris\"}".to_string(),
  1023. };
  1024. let encoded = serde_json::to_string(&delta).expect("delta should serialize");
  1025. let decoded: ContentBlockDelta =
  1026. serde_json::from_str(&encoded).expect("delta should deserialize");
  1027. assert_eq!(decoded, delta);
  1028. }
  1029. #[test]
  1030. fn request_id_uses_primary_or_fallback_header() {
  1031. let mut headers = reqwest::header::HeaderMap::new();
  1032. headers.insert(REQUEST_ID_HEADER, "req_primary".parse().expect("header"));
  1033. assert_eq!(
  1034. super::request_id_from_headers(&headers).as_deref(),
  1035. Some("req_primary")
  1036. );
  1037. headers.clear();
  1038. headers.insert(
  1039. ALT_REQUEST_ID_HEADER,
  1040. "req_fallback".parse().expect("header"),
  1041. );
  1042. assert_eq!(
  1043. super::request_id_from_headers(&headers).as_deref(),
  1044. Some("req_fallback")
  1045. );
  1046. }
  1047. #[test]
  1048. fn auth_source_applies_headers() {
  1049. let auth = AuthSource::ApiKeyAndBearer {
  1050. api_key: "test-key".to_string(),
  1051. bearer_token: "proxy-token".to_string(),
  1052. };
  1053. let request = auth
  1054. .apply(reqwest::Client::new().post("https://example.test"))
  1055. .build()
  1056. .expect("request build");
  1057. let headers = request.headers();
  1058. assert_eq!(
  1059. headers.get("x-api-key").and_then(|v| v.to_str().ok()),
  1060. Some("test-key")
  1061. );
  1062. assert_eq!(
  1063. headers.get("authorization").and_then(|v| v.to_str().ok()),
  1064. Some("Bearer proxy-token")
  1065. );
  1066. }
  1067. }
备用站点 当前处于降级运行的备用站点,仅供应急访问,数据和功能可能不是最新。