use std::collections::VecDeque;
use std::pin::Pin;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use serde::Deserialize;
use serde::de::Error as _;
use serde_json::Value;
use tokio_stream::{Stream, StreamExt};
use crate::config::{AuthKind, Config};
use crate::error::{ApiError, Error};
use crate::types::{
ContentBlock, ContentDelta, MessageDeltaUsage, MessageRequest, MessageResponse, StreamEvent,
Usage,
};
const ANTHROPIC_VERSION: &str = "2023-06-01";
const OAUTH_BETA: &str = "oauth-2025-04-20";
const BACKOFF_BASE_MS: u64 = 500;
const BACKOFF_MAX_MS: u64 = 30_000;
#[derive(Deserialize)]
struct ErrorEnvelope {
error: ErrorBody,
}
#[derive(Deserialize)]
struct ErrorBody {
#[serde(rename = "type")]
kind: String,
message: String,
}
#[derive(Debug, Clone)]
pub struct Client {
http: reqwest::Client,
config: Config,
url: String,
}
impl Client {
pub fn new(config: Config) -> Result<Self, Error> {
let http = reqwest::Client::builder()
.timeout(config.timeout)
.build()
.map_err(Error::Transport)?;
let url = format!("{}/v1/messages", config.base_url.trim_end_matches('/'));
Ok(Self { http, config, url })
}
pub fn from_env() -> Result<Self, Error> {
Self::new(Config::from_env())
}
#[must_use]
pub fn config(&self) -> &Config {
&self.config
}
pub async fn send_message(&self, request: &MessageRequest) -> Result<MessageResponse, Error> {
let mut attempt: u32 = 0;
loop {
match self.try_send(request).await {
Ok(response) => return Ok(response),
Err(err) => {
if attempt >= self.config.max_retries || !err.is_retryable() {
return Err(err);
}
let delay = err.retry_after().unwrap_or_else(|| backoff_delay(attempt));
tokio::time::sleep(delay).await;
attempt += 1;
}
}
}
}
pub async fn send_message_value(&self, request: &Value) -> Result<MessageResponse, Error> {
let mut attempt: u32 = 0;
loop {
match self.try_send_value(request).await {
Ok(response) => return Ok(response),
Err(err) => {
if attempt >= self.config.max_retries || !err.is_retryable() {
return Err(err);
}
let delay = err.retry_after().unwrap_or_else(|| backoff_delay(attempt));
tokio::time::sleep(delay).await;
attempt += 1;
}
}
}
}
pub async fn stream_message(&self, request: &MessageRequest) -> Result<MessageStream, Error> {
let mut streaming = request.clone();
streaming.stream = true;
let mut attempt: u32 = 0;
loop {
match self.try_stream(&streaming).await {
Ok(stream) => return Ok(stream),
Err(err) => {
if attempt >= self.config.max_retries || !err.is_retryable() {
return Err(err);
}
let delay = err.retry_after().unwrap_or_else(|| backoff_delay(attempt));
tokio::time::sleep(delay).await;
attempt += 1;
}
}
}
}
pub async fn stream_message_value(&self, request: &Value) -> Result<MessageStream, Error> {
let mut streaming = request.clone();
if let Value::Object(map) = &mut streaming {
map.insert("stream".to_owned(), Value::Bool(true));
}
let mut attempt: u32 = 0;
loop {
match self.try_stream_value(&streaming).await {
Ok(stream) => return Ok(stream),
Err(err) => {
if attempt >= self.config.max_retries || !err.is_retryable() {
return Err(err);
}
let delay = err.retry_after().unwrap_or_else(|| backoff_delay(attempt));
tokio::time::sleep(delay).await;
attempt += 1;
}
}
}
}
async fn try_send(&self, request: &MessageRequest) -> Result<MessageResponse, Error> {
let response = self
.post_builder()
.json(request)
.send()
.await
.map_err(Error::Transport)?;
let status = response.status();
let request_id = header_string(response.headers(), "request-id");
let retry_after = retry_after_of(response.headers());
let body = response.bytes().await.map_err(Error::Transport)?;
if status.is_success() {
return serde_json::from_slice(&body).map_err(Error::Decode);
}
Err(error_from_body(
status.as_u16(),
request_id,
retry_after,
&body,
))
}
async fn try_send_value(&self, request: &Value) -> Result<MessageResponse, Error> {
let response = self
.post_builder()
.json(request)
.send()
.await
.map_err(Error::Transport)?;
let status = response.status();
let request_id = header_string(response.headers(), "request-id");
let retry_after = retry_after_of(response.headers());
let body = response.bytes().await.map_err(Error::Transport)?;
if status.is_success() {
return serde_json::from_slice(&body).map_err(Error::Decode);
}
Err(error_from_body(
status.as_u16(),
request_id,
retry_after,
&body,
))
}
async fn try_stream(&self, request: &MessageRequest) -> Result<MessageStream, Error> {
let response = self
.post_builder()
.json(request)
.send()
.await
.map_err(Error::Transport)?;
let status = response.status();
let request_id = header_string(response.headers(), "request-id");
if status.is_success() {
let bytes = response
.bytes_stream()
.map(|chunk| chunk.map(|b| b.to_vec()));
return Ok(MessageStream::new(Box::pin(bytes), request_id));
}
let retry_after = retry_after_of(response.headers());
let body = response.bytes().await.map_err(Error::Transport)?;
Err(error_from_body(
status.as_u16(),
request_id,
retry_after,
&body,
))
}
async fn try_stream_value(&self, request: &Value) -> Result<MessageStream, Error> {
let response = self
.post_builder()
.json(request)
.send()
.await
.map_err(Error::Transport)?;
let status = response.status();
let request_id = header_string(response.headers(), "request-id");
if status.is_success() {
let bytes = response
.bytes_stream()
.map(|chunk| chunk.map(|b| b.to_vec()));
return Ok(MessageStream::new(Box::pin(bytes), request_id));
}
let retry_after = retry_after_of(response.headers());
let body = response.bytes().await.map_err(Error::Transport)?;
Err(error_from_body(
status.as_u16(),
request_id,
retry_after,
&body,
))
}
fn post_builder(&self) -> reqwest::RequestBuilder {
let mut builder = self
.http
.post(&self.url)
.header("anthropic-version", ANTHROPIC_VERSION);
if let Some(api_key) = &self.config.api_key {
match self.config.auth_kind {
AuthKind::ApiKey => {
builder = builder.header("x-api-key", api_key);
}
AuthKind::Bearer => {
builder = builder
.header("authorization", format!("Bearer {api_key}"))
.header("anthropic-beta", OAUTH_BETA);
}
}
}
builder
}
}
fn retry_after_of(headers: &reqwest::header::HeaderMap) -> Option<Duration> {
headers
.get("retry-after")
.and_then(|value| value.to_str().ok())
.and_then(|value| value.trim().parse::<u64>().ok())
.map(Duration::from_secs)
}
fn error_from_body(
status: u16,
request_id: Option<String>,
retry_after: Option<Duration>,
body: &[u8],
) -> Error {
match serde_json::from_slice::<ErrorEnvelope>(body) {
Ok(envelope) => Error::Api(ApiError {
status,
kind: envelope.error.kind,
message: envelope.error.message,
request_id,
retry_after,
}),
Err(_) => Error::Unexpected {
status,
body: String::from_utf8_lossy(body).into_owned(),
},
}
}
type ByteStream = Pin<Box<dyn Stream<Item = reqwest::Result<Vec<u8>>> + Send>>;
pub struct MessageStream {
bytes: ByteStream,
buffer: Vec<u8>,
data: String,
finished: bool,
request_id: Option<String>,
pending: VecDeque<Result<StreamEvent, Error>>,
}
impl MessageStream {
fn new(bytes: ByteStream, request_id: Option<String>) -> Self {
Self {
bytes,
buffer: Vec::new(),
data: String::new(),
finished: false,
request_id,
pending: VecDeque::new(),
}
}
#[must_use]
pub fn request_id(&self) -> Option<&str> {
self.request_id.as_deref()
}
pub async fn next_event(&mut self) -> Option<Result<StreamEvent, Error>> {
loop {
if let Some(event) = self.pending.pop_front() {
return Some(event);
}
if self.finished {
return None;
}
match self.bytes.next().await {
Some(Ok(chunk)) => {
self.buffer.extend_from_slice(&chunk);
self.drain_lines();
}
Some(Err(err)) => {
self.finished = true;
return Some(Err(Error::Transport(err)));
}
None => {
self.finished = true;
self.flush_tail();
}
}
}
}
pub async fn get_final_message(mut self) -> Result<MessageResponse, Error> {
let mut accumulator = MessageAccumulator::new();
while let Some(event) = self.next_event().await {
accumulator.apply(&event?)?;
}
accumulator.into_message()
}
fn drain_lines(&mut self) {
while let Some(pos) = self.buffer.iter().position(|&byte| byte == b'\n') {
let line: Vec<u8> = self.buffer.drain(..=pos).collect();
let mut line = &line[..line.len() - 1];
if line.last() == Some(&b'\r') {
line = &line[..line.len() - 1];
}
self.handle_line(line);
}
}
fn flush_tail(&mut self) {
if !self.buffer.is_empty() {
let line = std::mem::take(&mut self.buffer);
let mut line = &line[..];
if line.last() == Some(&b'\r') {
line = &line[..line.len() - 1];
}
self.handle_line(line);
}
self.dispatch();
}
fn handle_line(&mut self, line: &[u8]) {
if line.is_empty() {
self.dispatch();
return;
}
if line[0] == b':' {
return;
}
let (field, value) = split_field(line);
if field == b"data" {
if !self.data.is_empty() {
self.data.push('\n');
}
self.data.push_str(&String::from_utf8_lossy(value));
}
}
fn dispatch(&mut self) {
if self.data.is_empty() {
return;
}
let data = std::mem::take(&mut self.data);
let event = match serde_json::from_str::<Value>(&data) {
Ok(value) => StreamEvent::from_value(value),
Err(err) => Err(Error::Decode(err)),
};
self.pending.push_back(event);
}
}
fn split_field(line: &[u8]) -> (&[u8], &[u8]) {
match line.iter().position(|&byte| byte == b':') {
Some(index) => {
let field = &line[..index];
let mut value = &line[index + 1..];
if value.first() == Some(&b' ') {
value = &value[1..];
}
(field, value)
}
None => (line, &[]),
}
}
#[derive(Default)]
pub struct MessageAccumulator {
message: Option<MessageResponse>,
blocks: Vec<BlockState>,
}
struct BlockState {
block: ContentBlock,
partial_json: String,
}
impl BlockState {
fn apply_delta(&mut self, delta: &ContentDelta) {
match delta {
ContentDelta::Text { text } => {
if let ContentBlock::Text { text: current } = &mut self.block {
current.push_str(text);
}
}
ContentDelta::Thinking { thinking } => {
if let ContentBlock::Thinking {
thinking: current, ..
} = &mut self.block
{
current.push_str(thinking);
}
}
ContentDelta::Signature { signature } => {
if let ContentBlock::Thinking {
signature: current, ..
} = &mut self.block
{
match current {
Some(existing) => existing.push_str(signature),
None => *current = Some(signature.clone()),
}
}
}
ContentDelta::InputJson { partial_json } => {
if matches!(self.block, ContentBlock::ToolUse { .. }) {
self.partial_json.push_str(partial_json);
}
}
ContentDelta::Unknown(_) => {}
}
}
fn finalize(&mut self) -> Result<(), Error> {
if let ContentBlock::ToolUse { input, .. } = &mut self.block
&& !self.partial_json.is_empty()
{
*input = serde_json::from_str(&self.partial_json).map_err(Error::Decode)?;
}
Ok(())
}
}
impl MessageAccumulator {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn apply(&mut self, event: &StreamEvent) -> Result<(), Error> {
match event {
StreamEvent::MessageStart(message) => {
let mut message = message.clone();
message.content.clear();
self.message = Some(message);
self.blocks.clear();
}
StreamEvent::ContentBlockStart {
index,
content_block,
} => {
while self.blocks.len() <= *index {
self.blocks.push(BlockState {
block: ContentBlock::text(""),
partial_json: String::new(),
});
}
self.blocks[*index] = BlockState {
block: content_block.clone(),
partial_json: String::new(),
};
}
StreamEvent::ContentBlockDelta { index, delta } => {
if let Some(block) = self.blocks.get_mut(*index) {
block.apply_delta(delta);
}
}
StreamEvent::ContentBlockStop { index } => {
if let Some(block) = self.blocks.get_mut(*index) {
block.finalize()?;
}
}
StreamEvent::MessageDelta {
stop_reason,
stop_sequence,
usage,
} => {
if let Some(message) = self.message.as_mut() {
if stop_reason.is_some() {
message.stop_reason = stop_reason.clone();
}
if stop_sequence.is_some() {
message.stop_sequence = stop_sequence.clone();
}
merge_usage(&mut message.usage, usage);
}
}
StreamEvent::MessageStop | StreamEvent::Ping | StreamEvent::Unknown(_) => {}
}
Ok(())
}
pub fn into_message(self) -> Result<MessageResponse, Error> {
let mut message = self.message.ok_or_else(|| {
Error::Decode(serde_json::Error::custom(
"streaming response ended before a message_start event",
))
})?;
message.content = self.blocks.into_iter().map(|state| state.block).collect();
Ok(message)
}
}
fn merge_usage(usage: &mut Usage, delta: &MessageDeltaUsage) {
if let Some(input) = delta.input_tokens {
usage.input_tokens = input;
}
if let Some(output) = delta.output_tokens {
usage.output_tokens = output;
}
if delta.cache_creation_input_tokens.is_some() {
usage.cache_creation_input_tokens = delta.cache_creation_input_tokens;
}
if delta.cache_read_input_tokens.is_some() {
usage.cache_read_input_tokens = delta.cache_read_input_tokens;
}
}
fn header_string(headers: &reqwest::header::HeaderMap, name: &str) -> Option<String> {
headers
.get(name)
.and_then(|value| value.to_str().ok())
.map(str::to_owned)
}
fn backoff_delay(attempt: u32) -> Duration {
let window = BACKOFF_BASE_MS
.saturating_mul(1u64 << attempt.min(6))
.min(BACKOFF_MAX_MS);
let half = window / 2;
let jitter = (jitter_fraction() * half as f64) as u64;
Duration::from_millis(half + jitter)
}
fn jitter_fraction() -> f64 {
let nanos = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|elapsed| elapsed.subsec_nanos())
.unwrap_or(0);
f64::from(nanos % 1000) / 1000.0
}