client.rs 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424
  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)]
  15. pub struct AnthropicClient {
  16. http: reqwest::Client,
  17. api_key: String,
  18. auth_token: Option<String>,
  19. base_url: String,
  20. max_retries: u32,
  21. initial_backoff: Duration,
  22. max_backoff: Duration,
  23. }
  24. impl AnthropicClient {
  25. #[must_use]
  26. pub fn new(api_key: impl Into<String>) -> Self {
  27. Self {
  28. http: reqwest::Client::new(),
  29. api_key: api_key.into(),
  30. auth_token: None,
  31. base_url: DEFAULT_BASE_URL.to_string(),
  32. max_retries: DEFAULT_MAX_RETRIES,
  33. initial_backoff: DEFAULT_INITIAL_BACKOFF,
  34. max_backoff: DEFAULT_MAX_BACKOFF,
  35. }
  36. }
  37. pub fn from_env() -> Result<Self, ApiError> {
  38. Ok(Self::new(read_api_key()?)
  39. .with_auth_token(read_auth_token())
  40. .with_base_url(read_base_url()))
  41. }
  42. #[must_use]
  43. pub fn with_auth_token(mut self, auth_token: Option<String>) -> Self {
  44. self.auth_token = auth_token.filter(|token| !token.is_empty());
  45. self
  46. }
  47. #[must_use]
  48. pub fn with_base_url(mut self, base_url: impl Into<String>) -> Self {
  49. self.base_url = base_url.into();
  50. self
  51. }
  52. #[must_use]
  53. pub fn with_retry_policy(
  54. mut self,
  55. max_retries: u32,
  56. initial_backoff: Duration,
  57. max_backoff: Duration,
  58. ) -> Self {
  59. self.max_retries = max_retries;
  60. self.initial_backoff = initial_backoff;
  61. self.max_backoff = max_backoff;
  62. self
  63. }
  64. pub async fn send_message(
  65. &self,
  66. request: &MessageRequest,
  67. ) -> Result<MessageResponse, ApiError> {
  68. let request = MessageRequest {
  69. stream: false,
  70. ..request.clone()
  71. };
  72. let response = self.send_with_retry(&request).await?;
  73. let request_id = request_id_from_headers(response.headers());
  74. let mut response = response
  75. .json::<MessageResponse>()
  76. .await
  77. .map_err(ApiError::from)?;
  78. if response.request_id.is_none() {
  79. response.request_id = request_id;
  80. }
  81. Ok(response)
  82. }
  83. pub async fn stream_message(
  84. &self,
  85. request: &MessageRequest,
  86. ) -> Result<MessageStream, ApiError> {
  87. let response = self
  88. .send_with_retry(&request.clone().with_streaming())
  89. .await?;
  90. Ok(MessageStream {
  91. request_id: request_id_from_headers(response.headers()),
  92. response,
  93. parser: SseParser::new(),
  94. pending: VecDeque::new(),
  95. done: false,
  96. })
  97. }
  98. async fn send_with_retry(
  99. &self,
  100. request: &MessageRequest,
  101. ) -> Result<reqwest::Response, ApiError> {
  102. let mut attempts = 0;
  103. let mut last_error: Option<ApiError>;
  104. loop {
  105. attempts += 1;
  106. match self.send_raw_request(request).await {
  107. Ok(response) => match expect_success(response).await {
  108. Ok(response) => return Ok(response),
  109. Err(error) if error.is_retryable() && attempts <= self.max_retries + 1 => {
  110. last_error = Some(error);
  111. }
  112. Err(error) => return Err(error),
  113. },
  114. Err(error) if error.is_retryable() && attempts <= self.max_retries + 1 => {
  115. last_error = Some(error);
  116. }
  117. Err(error) => return Err(error),
  118. }
  119. if attempts > self.max_retries {
  120. break;
  121. }
  122. tokio::time::sleep(self.backoff_for_attempt(attempts)?).await;
  123. }
  124. Err(ApiError::RetriesExhausted {
  125. attempts,
  126. last_error: Box::new(last_error.expect("retry loop must capture an error")),
  127. })
  128. }
  129. async fn send_raw_request(
  130. &self,
  131. request: &MessageRequest,
  132. ) -> Result<reqwest::Response, ApiError> {
  133. let request_url = format!("{}/v1/messages", self.base_url.trim_end_matches('/'));
  134. let resolved_base_url = self.base_url.trim_end_matches('/');
  135. eprintln!("[anthropic-client] resolved_base_url={resolved_base_url}");
  136. eprintln!("[anthropic-client] request_url={request_url}");
  137. let mut request_builder = self
  138. .http
  139. .post(&request_url)
  140. .header("x-api-key", &self.api_key)
  141. .header("anthropic-version", ANTHROPIC_VERSION)
  142. .header("content-type", "application/json");
  143. let auth_header = self.auth_token.as_ref().map(|_| "Bearer [REDACTED]").unwrap_or("<absent>");
  144. eprintln!("[anthropic-client] headers x-api-key=[REDACTED] authorization={auth_header} anthropic-version={ANTHROPIC_VERSION} content-type=application/json");
  145. if let Some(auth_token) = &self.auth_token {
  146. request_builder = request_builder.bearer_auth(auth_token);
  147. }
  148. request_builder
  149. .json(request)
  150. .send()
  151. .await
  152. .map_err(ApiError::from)
  153. }
  154. fn backoff_for_attempt(&self, attempt: u32) -> Result<Duration, ApiError> {
  155. let Some(multiplier) = 1_u32.checked_shl(attempt.saturating_sub(1)) else {
  156. return Err(ApiError::BackoffOverflow {
  157. attempt,
  158. base_delay: self.initial_backoff,
  159. });
  160. };
  161. Ok(self
  162. .initial_backoff
  163. .checked_mul(multiplier)
  164. .map_or(self.max_backoff, |delay| delay.min(self.max_backoff)))
  165. }
  166. }
  167. fn read_api_key() -> Result<String, ApiError> {
  168. match std::env::var("ANTHROPIC_API_KEY") {
  169. Ok(api_key) if !api_key.is_empty() => Ok(api_key),
  170. Ok(_) => Err(ApiError::MissingApiKey),
  171. Err(std::env::VarError::NotPresent) => match std::env::var("ANTHROPIC_AUTH_TOKEN") {
  172. Ok(api_key) if !api_key.is_empty() => Ok(api_key),
  173. Ok(_) => Err(ApiError::MissingApiKey),
  174. Err(std::env::VarError::NotPresent) => Err(ApiError::MissingApiKey),
  175. Err(error) => Err(ApiError::from(error)),
  176. },
  177. Err(error) => Err(ApiError::from(error)),
  178. }
  179. }
  180. fn read_auth_token() -> Option<String> {
  181. match std::env::var("ANTHROPIC_AUTH_TOKEN") {
  182. Ok(token) if !token.is_empty() => Some(token),
  183. _ => None,
  184. }
  185. }
  186. fn read_base_url() -> String {
  187. std::env::var("ANTHROPIC_BASE_URL").unwrap_or_else(|_| DEFAULT_BASE_URL.to_string())
  188. }
  189. fn request_id_from_headers(headers: &reqwest::header::HeaderMap) -> Option<String> {
  190. headers
  191. .get(REQUEST_ID_HEADER)
  192. .or_else(|| headers.get(ALT_REQUEST_ID_HEADER))
  193. .and_then(|value| value.to_str().ok())
  194. .map(ToOwned::to_owned)
  195. }
  196. #[derive(Debug)]
  197. pub struct MessageStream {
  198. request_id: Option<String>,
  199. response: reqwest::Response,
  200. parser: SseParser,
  201. pending: VecDeque<StreamEvent>,
  202. done: bool,
  203. }
  204. impl MessageStream {
  205. #[must_use]
  206. pub fn request_id(&self) -> Option<&str> {
  207. self.request_id.as_deref()
  208. }
  209. pub async fn next_event(&mut self) -> Result<Option<StreamEvent>, ApiError> {
  210. loop {
  211. if let Some(event) = self.pending.pop_front() {
  212. return Ok(Some(event));
  213. }
  214. if self.done {
  215. let remaining = self.parser.finish()?;
  216. self.pending.extend(remaining);
  217. if let Some(event) = self.pending.pop_front() {
  218. return Ok(Some(event));
  219. }
  220. return Ok(None);
  221. }
  222. match self.response.chunk().await? {
  223. Some(chunk) => {
  224. self.pending.extend(self.parser.push(&chunk)?);
  225. }
  226. None => {
  227. self.done = true;
  228. }
  229. }
  230. }
  231. }
  232. }
  233. async fn expect_success(response: reqwest::Response) -> Result<reqwest::Response, ApiError> {
  234. let status = response.status();
  235. if status.is_success() {
  236. return Ok(response);
  237. }
  238. let body = response.text().await.unwrap_or_else(|_| String::new());
  239. let parsed_error = serde_json::from_str::<AnthropicErrorEnvelope>(&body).ok();
  240. let retryable = is_retryable_status(status);
  241. Err(ApiError::Api {
  242. status,
  243. error_type: parsed_error
  244. .as_ref()
  245. .map(|error| error.error.error_type.clone()),
  246. message: parsed_error
  247. .as_ref()
  248. .map(|error| error.error.message.clone()),
  249. body,
  250. retryable,
  251. })
  252. }
  253. const fn is_retryable_status(status: reqwest::StatusCode) -> bool {
  254. matches!(status.as_u16(), 408 | 409 | 429 | 500 | 502 | 503 | 504)
  255. }
  256. #[derive(Debug, Deserialize)]
  257. struct AnthropicErrorEnvelope {
  258. error: AnthropicErrorBody,
  259. }
  260. #[derive(Debug, Deserialize)]
  261. struct AnthropicErrorBody {
  262. #[serde(rename = "type")]
  263. error_type: String,
  264. message: String,
  265. }
  266. #[cfg(test)]
  267. mod tests {
  268. use super::{ALT_REQUEST_ID_HEADER, REQUEST_ID_HEADER};
  269. use std::time::Duration;
  270. use crate::types::{ContentBlockDelta, MessageRequest};
  271. #[test]
  272. fn read_api_key_requires_presence() {
  273. std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
  274. std::env::remove_var("ANTHROPIC_API_KEY");
  275. let error = super::read_api_key().expect_err("missing key should error");
  276. assert!(matches!(error, crate::error::ApiError::MissingApiKey));
  277. }
  278. #[test]
  279. fn read_api_key_requires_non_empty_value() {
  280. std::env::set_var("ANTHROPIC_AUTH_TOKEN", "");
  281. std::env::remove_var("ANTHROPIC_API_KEY");
  282. let error = super::read_api_key().expect_err("empty key should error");
  283. assert!(matches!(error, crate::error::ApiError::MissingApiKey));
  284. }
  285. #[test]
  286. fn read_api_key_prefers_api_key_env() {
  287. std::env::set_var("ANTHROPIC_AUTH_TOKEN", "auth-token");
  288. std::env::set_var("ANTHROPIC_API_KEY", "legacy-key");
  289. assert_eq!(
  290. super::read_api_key().expect("api key should load"),
  291. "legacy-key"
  292. );
  293. std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
  294. std::env::remove_var("ANTHROPIC_API_KEY");
  295. }
  296. #[test]
  297. fn read_auth_token_reads_auth_token_env() {
  298. std::env::set_var("ANTHROPIC_AUTH_TOKEN", "auth-token");
  299. assert_eq!(super::read_auth_token().as_deref(), Some("auth-token"));
  300. std::env::remove_var("ANTHROPIC_AUTH_TOKEN");
  301. }
  302. #[test]
  303. fn message_request_stream_helper_sets_stream_true() {
  304. let request = MessageRequest {
  305. model: "claude-3-7-sonnet-latest".to_string(),
  306. max_tokens: 64,
  307. messages: vec![],
  308. system: None,
  309. tools: None,
  310. tool_choice: None,
  311. stream: false,
  312. };
  313. assert!(request.with_streaming().stream);
  314. }
  315. #[test]
  316. fn backoff_doubles_until_maximum() {
  317. let client = super::AnthropicClient::new("test-key").with_retry_policy(
  318. 3,
  319. Duration::from_millis(10),
  320. Duration::from_millis(25),
  321. );
  322. assert_eq!(
  323. client.backoff_for_attempt(1).expect("attempt 1"),
  324. Duration::from_millis(10)
  325. );
  326. assert_eq!(
  327. client.backoff_for_attempt(2).expect("attempt 2"),
  328. Duration::from_millis(20)
  329. );
  330. assert_eq!(
  331. client.backoff_for_attempt(3).expect("attempt 3"),
  332. Duration::from_millis(25)
  333. );
  334. }
  335. #[test]
  336. fn retryable_statuses_are_detected() {
  337. assert!(super::is_retryable_status(
  338. reqwest::StatusCode::TOO_MANY_REQUESTS
  339. ));
  340. assert!(super::is_retryable_status(
  341. reqwest::StatusCode::INTERNAL_SERVER_ERROR
  342. ));
  343. assert!(!super::is_retryable_status(
  344. reqwest::StatusCode::UNAUTHORIZED
  345. ));
  346. }
  347. #[test]
  348. fn tool_delta_variant_round_trips() {
  349. let delta = ContentBlockDelta::InputJsonDelta {
  350. partial_json: "{\"city\":\"Paris\"}".to_string(),
  351. };
  352. let encoded = serde_json::to_string(&delta).expect("delta should serialize");
  353. let decoded: ContentBlockDelta =
  354. serde_json::from_str(&encoded).expect("delta should deserialize");
  355. assert_eq!(decoded, delta);
  356. }
  357. #[test]
  358. fn request_id_uses_primary_or_fallback_header() {
  359. let mut headers = reqwest::header::HeaderMap::new();
  360. headers.insert(REQUEST_ID_HEADER, "req_primary".parse().expect("header"));
  361. assert_eq!(
  362. super::request_id_from_headers(&headers).as_deref(),
  363. Some("req_primary")
  364. );
  365. headers.clear();
  366. headers.insert(
  367. ALT_REQUEST_ID_HEADER,
  368. "req_fallback".parse().expect("header"),
  369. );
  370. assert_eq!(
  371. super::request_id_from_headers(&headers).as_deref(),
  372. Some("req_fallback")
  373. );
  374. }
  375. }
备用站点 当前处于降级运行的备用站点,仅供应急访问,数据和功能可能不是最新。