client.rs 18 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603
  1. use std::collections::VecDeque;
  2. use std::time::Duration;
  3. use serde::Deserialize;
  4. use crate::error::ApiError;
  5. use crate::sse::SseParser;
  6. use crate::types::{MessageRequest, MessageResponse, StreamEvent};
  7. const DEFAULT_BASE_URL: &str = "https://api.anthropic.com";
  8. const ANTHROPIC_VERSION: &str = "2023-06-01";
  9. const REQUEST_ID_HEADER: &str = "request-id";
  10. const ALT_REQUEST_ID_HEADER: &str = "x-request-id";
  11. const DEFAULT_INITIAL_BACKOFF: Duration = Duration::from_millis(200);
  12. const DEFAULT_MAX_BACKOFF: Duration = Duration::from_secs(2);
  13. const DEFAULT_MAX_RETRIES: u32 = 2;
  14. #[derive(Debug, Clone, PartialEq, Eq)]
  15. pub enum AuthSource {
  16. None,
  17. ApiKey(String),
  18. BearerToken(String),
  19. ApiKeyAndBearer {
  20. api_key: String,
  21. bearer_token: String,
  22. },
  23. }
  24. impl AuthSource {
  25. pub fn from_env() -> Result<Self, ApiError> {
  26. let api_key = read_env_non_empty("ANTHROPIC_API_KEY")?;
  27. let auth_token = read_env_non_empty("ANTHROPIC_AUTH_TOKEN")?;
  28. match (api_key, auth_token) {
  29. (Some(api_key), Some(bearer_token)) => Ok(Self::ApiKeyAndBearer {
  30. api_key,
  31. bearer_token,
  32. }),
  33. (Some(api_key), None) => Ok(Self::ApiKey(api_key)),
  34. (None, Some(bearer_token)) => Ok(Self::BearerToken(bearer_token)),
  35. (None, None) => Err(ApiError::MissingApiKey),
  36. }
  37. }
  38. #[must_use]
  39. pub fn api_key(&self) -> Option<&str> {
  40. match self {
  41. Self::ApiKey(api_key) | Self::ApiKeyAndBearer { api_key, .. } => Some(api_key),
  42. Self::None | Self::BearerToken(_) => None,
  43. }
  44. }
  45. #[must_use]
  46. pub fn bearer_token(&self) -> Option<&str> {
  47. match self {
  48. Self::BearerToken(token)
  49. | Self::ApiKeyAndBearer {
  50. bearer_token: token,
  51. ..
  52. } => Some(token),
  53. Self::None | Self::ApiKey(_) => None,
  54. }
  55. }
  56. #[must_use]
  57. pub fn masked_authorization_header(&self) -> &'static str {
  58. if self.bearer_token().is_some() {
  59. "Bearer [REDACTED]"
  60. } else {
  61. "<absent>"
  62. }
  63. }
  64. pub fn apply(&self, mut request_builder: reqwest::RequestBuilder) -> reqwest::RequestBuilder {
  65. if let Some(api_key) = self.api_key() {
  66. request_builder = request_builder.header("x-api-key", api_key);
  67. }
  68. if let Some(token) = self.bearer_token() {
  69. request_builder = request_builder.bearer_auth(token);
  70. }
  71. request_builder
  72. }
  73. }
  74. #[derive(Debug, Clone, PartialEq, Eq)]
  75. pub struct OAuthTokenSet {
  76. pub access_token: String,
  77. pub refresh_token: Option<String>,
  78. pub expires_at: Option<u64>,
  79. pub scopes: Vec<String>,
  80. }
  81. impl From<OAuthTokenSet> for AuthSource {
  82. fn from(value: OAuthTokenSet) -> Self {
  83. Self::BearerToken(value.access_token)
  84. }
  85. }
  86. #[derive(Debug, Clone)]
  87. pub struct AnthropicClient {
  88. http: reqwest::Client,
  89. auth: AuthSource,
  90. base_url: String,
  91. max_retries: u32,
  92. initial_backoff: Duration,
  93. max_backoff: Duration,
  94. }
  95. impl AnthropicClient {
  96. #[must_use]
  97. pub fn new(api_key: impl Into<String>) -> Self {
  98. Self {
  99. http: reqwest::Client::new(),
  100. auth: AuthSource::ApiKey(api_key.into()),
  101. base_url: DEFAULT_BASE_URL.to_string(),
  102. max_retries: DEFAULT_MAX_RETRIES,
  103. initial_backoff: DEFAULT_INITIAL_BACKOFF,
  104. max_backoff: DEFAULT_MAX_BACKOFF,
  105. }
  106. }
  107. #[must_use]
  108. pub fn from_auth(auth: AuthSource) -> Self {
  109. Self {
  110. http: reqwest::Client::new(),
  111. auth,
  112. base_url: DEFAULT_BASE_URL.to_string(),
  113. max_retries: DEFAULT_MAX_RETRIES,
  114. initial_backoff: DEFAULT_INITIAL_BACKOFF,
  115. max_backoff: DEFAULT_MAX_BACKOFF,
  116. }
  117. }
  118. pub fn from_env() -> Result<Self, ApiError> {
  119. Ok(Self::from_auth(AuthSource::from_env()?).with_base_url(read_base_url()))
  120. }
  121. #[must_use]
  122. pub fn with_auth_source(mut self, auth: AuthSource) -> Self {
  123. self.auth = auth;
  124. self
  125. }
  126. #[must_use]
  127. pub fn with_auth_token(mut self, auth_token: Option<String>) -> Self {
  128. match (
  129. self.auth.api_key().map(ToOwned::to_owned),
  130. auth_token.filter(|token| !token.is_empty()),
  131. ) {
  132. (Some(api_key), Some(bearer_token)) => {
  133. self.auth = AuthSource::ApiKeyAndBearer {
  134. api_key,
  135. bearer_token,
  136. };
  137. }
  138. (Some(api_key), None) => {
  139. self.auth = AuthSource::ApiKey(api_key);
  140. }
  141. (None, Some(bearer_token)) => {
  142. self.auth = AuthSource::BearerToken(bearer_token);
  143. }
  144. (None, None) => {
  145. self.auth = AuthSource::None;
  146. }
  147. }
  148. self
  149. }
  150. #[must_use]
  151. pub fn with_base_url(mut self, base_url: impl Into<String>) -> Self {
  152. self.base_url = base_url.into();
  153. self
  154. }
  155. #[must_use]
  156. pub fn with_retry_policy(
  157. mut self,
  158. max_retries: u32,
  159. initial_backoff: Duration,
  160. max_backoff: Duration,
  161. ) -> Self {
  162. self.max_retries = max_retries;
  163. self.initial_backoff = initial_backoff;
  164. self.max_backoff = max_backoff;
  165. self
  166. }
  167. #[must_use]
  168. pub fn auth_source(&self) -> &AuthSource {
  169. &self.auth
  170. }
  171. pub async fn send_message(
  172. &self,
  173. request: &MessageRequest,
  174. ) -> Result<MessageResponse, ApiError> {
  175. let request = MessageRequest {
  176. stream: false,
  177. ..request.clone()
  178. };
  179. let response = self.send_with_retry(&request).await?;
  180. let request_id = request_id_from_headers(response.headers());
  181. let mut response = response
  182. .json::<MessageResponse>()
  183. .await
  184. .map_err(ApiError::from)?;
  185. if response.request_id.is_none() {
  186. response.request_id = request_id;
  187. }
  188. Ok(response)
  189. }
  190. pub async fn stream_message(
  191. &self,
  192. request: &MessageRequest,
  193. ) -> Result<MessageStream, ApiError> {
  194. let response = self
  195. .send_with_retry(&request.clone().with_streaming())
  196. .await?;
  197. Ok(MessageStream {
  198. request_id: request_id_from_headers(response.headers()),
  199. response,
  200. parser: SseParser::new(),
  201. pending: VecDeque::new(),
  202. done: false,
  203. })
  204. }
  205. async fn send_with_retry(
  206. &self,
  207. request: &MessageRequest,
  208. ) -> Result<reqwest::Response, ApiError> {
  209. let mut attempts = 0;
  210. let mut last_error: Option<ApiError>;
  211. loop {
  212. attempts += 1;
  213. match self.send_raw_request(request).await {
  214. Ok(response) => match expect_success(response).await {
  215. Ok(response) => return Ok(response),
  216. Err(error) if error.is_retryable() && attempts <= self.max_retries + 1 => {
  217. last_error = Some(error);
  218. }
  219. Err(error) => return Err(error),
  220. },
  221. Err(error) if error.is_retryable() && attempts <= self.max_retries + 1 => {
  222. last_error = Some(error);
  223. }
  224. Err(error) => return Err(error),
  225. }
  226. if attempts > self.max_retries {
  227. break;
  228. }
  229. tokio::time::sleep(self.backoff_for_attempt(attempts)?).await;
  230. }
  231. Err(ApiError::RetriesExhausted {
  232. attempts,
  233. last_error: Box::new(last_error.expect("retry loop must capture an error")),
  234. })
  235. }
  236. async fn send_raw_request(
  237. &self,
  238. request: &MessageRequest,
  239. ) -> Result<reqwest::Response, ApiError> {
  240. let request_url = format!("{}/v1/messages", self.base_url.trim_end_matches('/'));
  241. let resolved_base_url = self.base_url.trim_end_matches('/');
  242. eprintln!("[anthropic-client] resolved_base_url={resolved_base_url}");
  243. eprintln!("[anthropic-client] request_url={request_url}");
  244. let request_builder = self
  245. .http
  246. .post(&request_url)
  247. .header("anthropic-version", ANTHROPIC_VERSION)
  248. .header("content-type", "application/json");
  249. let mut request_builder = self.auth.apply(request_builder);
  250. eprintln!(
  251. "[anthropic-client] headers x-api-key={} authorization={} anthropic-version={ANTHROPIC_VERSION} content-type=application/json",
  252. if self.auth.api_key().is_some() {
  253. "[REDACTED]"
  254. } else {
  255. "<absent>"
  256. },
  257. self.auth.masked_authorization_header()
  258. );
  259. request_builder = request_builder.json(request);
  260. request_builder.send().await.map_err(ApiError::from)
  261. }
  262. fn backoff_for_attempt(&self, attempt: u32) -> Result<Duration, ApiError> {
  263. let Some(multiplier) = 1_u32.checked_shl(attempt.saturating_sub(1)) else {
  264. return Err(ApiError::BackoffOverflow {
  265. attempt,
  266. base_delay: self.initial_backoff,
  267. });
  268. };
  269. Ok(self
  270. .initial_backoff
  271. .checked_mul(multiplier)
  272. .map_or(self.max_backoff, |delay| delay.min(self.max_backoff)))
  273. }
  274. }
  275. fn read_env_non_empty(key: &str) -> Result<Option<String>, ApiError> {
  276. match std::env::var(key) {
  277. Ok(value) if !value.is_empty() => Ok(Some(value)),
  278. Ok(_) | Err(std::env::VarError::NotPresent) => Ok(None),
  279. Err(error) => Err(ApiError::from(error)),
  280. }
  281. }
  282. #[cfg(test)]
  283. fn read_api_key() -> Result<String, ApiError> {
  284. let auth = AuthSource::from_env()?;
  285. auth.api_key()
  286. .or_else(|| auth.bearer_token())
  287. .map(ToOwned::to_owned)
  288. .ok_or(ApiError::MissingApiKey)
  289. }
  290. #[cfg(test)]
  291. fn read_auth_token() -> Option<String> {
  292. read_env_non_empty("ANTHROPIC_AUTH_TOKEN")
  293. .ok()
  294. .and_then(std::convert::identity)
  295. }
  296. fn read_base_url() -> String {
  297. std::env::var("ANTHROPIC_BASE_URL").unwrap_or_else(|_| DEFAULT_BASE_URL.to_string())
  298. }
  299. fn request_id_from_headers(headers: &reqwest::header::HeaderMap) -> Option<String> {
  300. headers
  301. .get(REQUEST_ID_HEADER)
  302. .or_else(|| headers.get(ALT_REQUEST_ID_HEADER))
  303. .and_then(|value| value.to_str().ok())
  304. .map(ToOwned::to_owned)
  305. }
  306. #[derive(Debug)]
  307. pub struct MessageStream {
  308. request_id: Option<String>,
  309. response: reqwest::Response,
  310. parser: SseParser,
  311. pending: VecDeque<StreamEvent>,
  312. done: bool,
  313. }
  314. impl MessageStream {
  315. #[must_use]
  316. pub fn request_id(&self) -> Option<&str> {
  317. self.request_id.as_deref()
  318. }
  319. pub async fn next_event(&mut self) -> Result<Option<StreamEvent>, ApiError> {
  320. loop {
  321. if let Some(event) = self.pending.pop_front() {
  322. return Ok(Some(event));
  323. }
  324. if self.done {
  325. let remaining = self.parser.finish()?;
  326. self.pending.extend(remaining);
  327. if let Some(event) = self.pending.pop_front() {
  328. return Ok(Some(event));
  329. }
  330. return Ok(None);
  331. }
  332. match self.response.chunk().await? {
  333. Some(chunk) => {
  334. self.pending.extend(self.parser.push(&chunk)?);
  335. }
  336. None => {
  337. self.done = true;
  338. }
  339. }
  340. }
  341. }
  342. }
  343. async fn expect_success(response: reqwest::Response) -> Result<reqwest::Response, ApiError> {
  344. let status = response.status();
  345. if status.is_success() {
  346. return Ok(response);
  347. }
  348. let body = response.text().await.unwrap_or_else(|_| String::new());
  349. let parsed_error = serde_json::from_str::<AnthropicErrorEnvelope>(&body).ok();
  350. let retryable = is_retryable_status(status);
  351. Err(ApiError::Api {
  352. status,
  353. error_type: parsed_error
  354. .as_ref()
  355. .map(|error| error.error.error_type.clone()),
  356. message: parsed_error
  357. .as_ref()
  358. .map(|error| error.error.message.clone()),
  359. body,
  360. retryable,
  361. })
  362. }
  363. const fn is_retryable_status(status: reqwest::StatusCode) -> bool {
  364. matches!(status.as_u16(), 408 | 409 | 429 | 500 | 502 | 503 | 504)
  365. }
  366. #[derive(Debug, Deserialize)]
  367. struct AnthropicErrorEnvelope {
  368. error: AnthropicErrorBody,
  369. }
  370. #[derive(Debug, Deserialize)]
  371. struct AnthropicErrorBody {
  372. #[serde(rename = "type")]
  373. error_type: String,
  374. message: String,
  375. }
  376. #[cfg(test)]
  377. mod tests {
  378. use super::{ALT_REQUEST_ID_HEADER, REQUEST_ID_HEADER};
  379. use std::sync::{Mutex, OnceLock};
  380. use std::time::Duration;
  381. use crate::client::{AuthSource, OAuthTokenSet};
  382. use crate::types::{ContentBlockDelta, MessageRequest};
  383. fn env_lock() -> std::sync::MutexGuard<'static, ()> {
  384. static LOCK: OnceLock<Mutex<()>> = OnceLock::new();
  385. LOCK.get_or_init(|| Mutex::new(()))
  386. .lock()
  387. .expect("env lock")
  388. }
  389. #[test]
  390. fn read_api_key_requires_presence() {
  391. let _guard = env_lock();
  392. std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
  393. std::env::remove_var("ANTHROPIC_API_KEY");
  394. let error = super::read_api_key().expect_err("missing key should error");
  395. assert!(matches!(error, crate::error::ApiError::MissingApiKey));
  396. }
  397. #[test]
  398. fn read_api_key_requires_non_empty_value() {
  399. let _guard = env_lock();
  400. std::env::set_var("ANTHROPIC_AUTH_TOKEN", "");
  401. std::env::remove_var("ANTHROPIC_API_KEY");
  402. let error = super::read_api_key().expect_err("empty key should error");
  403. assert!(matches!(error, crate::error::ApiError::MissingApiKey));
  404. }
  405. #[test]
  406. fn read_api_key_prefers_api_key_env() {
  407. let _guard = env_lock();
  408. std::env::set_var("ANTHROPIC_AUTH_TOKEN", "auth-token");
  409. std::env::set_var("ANTHROPIC_API_KEY", "legacy-key");
  410. assert_eq!(
  411. super::read_api_key().expect("api key should load"),
  412. "legacy-key"
  413. );
  414. std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
  415. std::env::remove_var("ANTHROPIC_API_KEY");
  416. }
  417. #[test]
  418. fn read_auth_token_reads_auth_token_env() {
  419. let _guard = env_lock();
  420. std::env::set_var("ANTHROPIC_AUTH_TOKEN", "auth-token");
  421. assert_eq!(super::read_auth_token().as_deref(), Some("auth-token"));
  422. std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
  423. }
  424. #[test]
  425. fn oauth_token_maps_to_bearer_auth_source() {
  426. let auth = AuthSource::from(OAuthTokenSet {
  427. access_token: "access-token".to_string(),
  428. refresh_token: Some("refresh".to_string()),
  429. expires_at: Some(123),
  430. scopes: vec!["scope:a".to_string()],
  431. });
  432. assert_eq!(auth.bearer_token(), Some("access-token"));
  433. assert_eq!(auth.api_key(), None);
  434. }
  435. #[test]
  436. fn auth_source_from_env_combines_api_key_and_bearer_token() {
  437. let _guard = env_lock();
  438. std::env::set_var("ANTHROPIC_AUTH_TOKEN", "auth-token");
  439. std::env::set_var("ANTHROPIC_API_KEY", "legacy-key");
  440. let auth = AuthSource::from_env().expect("env auth");
  441. assert_eq!(auth.api_key(), Some("legacy-key"));
  442. assert_eq!(auth.bearer_token(), Some("auth-token"));
  443. std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
  444. std::env::remove_var("ANTHROPIC_API_KEY");
  445. }
  446. #[test]
  447. fn message_request_stream_helper_sets_stream_true() {
  448. let request = MessageRequest {
  449. model: "claude-3-7-sonnet-latest".to_string(),
  450. max_tokens: 64,
  451. messages: vec![],
  452. system: None,
  453. tools: None,
  454. tool_choice: None,
  455. stream: false,
  456. };
  457. assert!(request.with_streaming().stream);
  458. }
  459. #[test]
  460. fn backoff_doubles_until_maximum() {
  461. let client = super::AnthropicClient::new("test-key").with_retry_policy(
  462. 3,
  463. Duration::from_millis(10),
  464. Duration::from_millis(25),
  465. );
  466. assert_eq!(
  467. client.backoff_for_attempt(1).expect("attempt 1"),
  468. Duration::from_millis(10)
  469. );
  470. assert_eq!(
  471. client.backoff_for_attempt(2).expect("attempt 2"),
  472. Duration::from_millis(20)
  473. );
  474. assert_eq!(
  475. client.backoff_for_attempt(3).expect("attempt 3"),
  476. Duration::from_millis(25)
  477. );
  478. }
  479. #[test]
  480. fn retryable_statuses_are_detected() {
  481. assert!(super::is_retryable_status(
  482. reqwest::StatusCode::TOO_MANY_REQUESTS
  483. ));
  484. assert!(super::is_retryable_status(
  485. reqwest::StatusCode::INTERNAL_SERVER_ERROR
  486. ));
  487. assert!(!super::is_retryable_status(
  488. reqwest::StatusCode::UNAUTHORIZED
  489. ));
  490. }
  491. #[test]
  492. fn tool_delta_variant_round_trips() {
  493. let delta = ContentBlockDelta::InputJsonDelta {
  494. partial_json: "{\"city\":\"Paris\"}".to_string(),
  495. };
  496. let encoded = serde_json::to_string(&delta).expect("delta should serialize");
  497. let decoded: ContentBlockDelta =
  498. serde_json::from_str(&encoded).expect("delta should deserialize");
  499. assert_eq!(decoded, delta);
  500. }
  501. #[test]
  502. fn request_id_uses_primary_or_fallback_header() {
  503. let mut headers = reqwest::header::HeaderMap::new();
  504. headers.insert(REQUEST_ID_HEADER, "req_primary".parse().expect("header"));
  505. assert_eq!(
  506. super::request_id_from_headers(&headers).as_deref(),
  507. Some("req_primary")
  508. );
  509. headers.clear();
  510. headers.insert(
  511. ALT_REQUEST_ID_HEADER,
  512. "req_fallback".parse().expect("header"),
  513. );
  514. assert_eq!(
  515. super::request_id_from_headers(&headers).as_deref(),
  516. Some("req_fallback")
  517. );
  518. }
  519. #[test]
  520. fn auth_source_applies_headers() {
  521. let auth = AuthSource::ApiKeyAndBearer {
  522. api_key: "test-key".to_string(),
  523. bearer_token: "proxy-token".to_string(),
  524. };
  525. let request = auth
  526. .apply(reqwest::Client::new().post("https://example.test"))
  527. .build()
  528. .expect("request build");
  529. let headers = request.headers();
  530. assert_eq!(
  531. headers.get("x-api-key").and_then(|v| v.to_str().ok()),
  532. Some("test-key")
  533. );
  534. assert_eq!(
  535. headers.get("authorization").and_then(|v| v.to_str().ok()),
  536. Some("Bearer proxy-token")
  537. );
  538. }
  539. }
备用站点 当前处于降级运行的备用站点,仅供应急访问,数据和功能可能不是最新。