use std::path::PathBuf;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, OnceLock};
use std::time::Duration;
use reqwest::header::{AUTHORIZATION, CONTENT_TYPE};
use reqwest::Method;
use serde_json::{Map, Value};
use crate::client::USER_AGENT;
use crate::discovery::env_var;
use crate::error::{code_for_status, Result, WritError};
use crate::models::{CrawlJob, CrawlStartParams};
const DEFAULT_CLOUD_URL: &str = "https://api.usewrit.app";
const CLIENT_ID_HEADER: &str = "X-Writ-Client-Id";
const DEFAULT_TIMEOUT: Duration = Duration::from_secs(30);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CloudTier {
Keyless,
Metered,
}
impl CloudTier {
pub fn as_str(&self) -> &'static str {
match self {
CloudTier::Keyless => "keyless",
CloudTier::Metered => "metered",
}
}
}
impl std::fmt::Display for CloudTier {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Debug, Clone)]
pub struct KeylessQuota {
pub tier: CloudTier,
pub requests_remaining: i64,
pub pages_remaining: i64,
pub requests_per_day: i64,
pub pages_per_day: i64,
pub reset_at: String,
pub upgrade_url: Option<String>,
}
#[derive(Debug, Clone)]
pub struct ScrapeResult {
pub url: String,
pub title: Option<String>,
pub format: String,
pub markdown: String,
pub counts: Map<String, Value>,
pub tier: CloudTier,
pub quota: Option<KeylessQuota>,
}
#[derive(Debug, Clone)]
pub struct MapEntry {
pub url: String,
pub score: f64,
pub title: Option<String>,
}
#[derive(Debug, Clone, Default)]
pub struct MapCounts {
pub returned: i64,
pub total: i64,
}
#[derive(Debug, Clone)]
pub struct MapResult {
pub url: String,
pub host: Option<String>,
pub urls: Vec<MapEntry>,
pub counts: MapCounts,
pub tier: CloudTier,
pub quota: Option<KeylessQuota>,
}
#[derive(Debug, Clone, Default)]
pub struct MapOptions {
pub search: Option<String>,
pub limit: Option<i64>,
}
#[derive(Debug, Default, Clone)]
pub struct CloudClientBuilder {
api_key: Option<String>,
cloud_url: Option<String>,
client_id: Option<String>,
timeout: Option<Duration>,
}
impl CloudClientBuilder {
pub fn api_key(mut self, api_key: impl Into<String>) -> Self {
self.api_key = Some(api_key.into());
self
}
pub fn cloud_url(mut self, cloud_url: impl Into<String>) -> Self {
self.cloud_url = Some(cloud_url.into());
self
}
pub fn client_id(mut self, client_id: impl Into<String>) -> Self {
self.client_id = Some(client_id.into());
self
}
pub fn timeout(mut self, timeout: Duration) -> Self {
self.timeout = Some(timeout);
self
}
pub fn build(self) -> Result<CloudClient> {
let api_key = self.api_key.or_else(|| env_var("WRIT_API_KEY"));
let base = self
.cloud_url
.or_else(|| env_var("WRIT_CLOUD_URL"))
.unwrap_or_else(|| DEFAULT_CLOUD_URL.to_string())
.trim_end_matches('/')
.to_string();
let client_id_override = self.client_id.or_else(|| env_var("WRIT_CLIENT_ID"));
let http = reqwest::Client::builder()
.timeout(self.timeout.unwrap_or(DEFAULT_TIMEOUT))
.user_agent(USER_AGENT)
.build()
.map_err(|e| WritError::Connection(format!("building cloud http client: {e}")))?;
Ok(CloudClient {
inner: Arc::new(CloudInner {
api_key,
base,
client_id_override,
client_id_cache: OnceLock::new(),
http,
}),
})
}
}
#[derive(Debug, Clone)]
pub struct CloudClient {
inner: Arc<CloudInner>,
}
#[derive(Debug)]
struct CloudInner {
api_key: Option<String>,
base: String,
client_id_override: Option<String>,
client_id_cache: OnceLock<String>,
http: reqwest::Client,
}
impl CloudClient {
pub fn builder() -> CloudClientBuilder {
CloudClientBuilder::default()
}
pub fn from_env() -> Result<CloudClient> {
CloudClientBuilder::default().build()
}
pub fn base_url(&self) -> &str {
&self.inner.base
}
pub fn tier(&self) -> CloudTier {
if self.inner.api_key.is_some() {
CloudTier::Metered
} else {
CloudTier::Keyless
}
}
pub async fn scrape(&self, url: &str) -> Result<ScrapeResult> {
let path = if self.inner.api_key.is_some() {
"/api/crawl/scrape"
} else {
"/v1/keyless/scrape"
};
let raw = self
.send(Method::POST, path, Some(&serde_json::json!({ "url": url })))
.await?;
Ok(normalize_scrape(&raw, self.tier()))
}
pub async fn map(&self, url: &str, opts: &MapOptions) -> Result<MapResult> {
let path = if self.inner.api_key.is_some() {
"/api/crawl/map"
} else {
"/v1/keyless/map"
};
let mut body = Map::new();
body.insert("url".into(), Value::String(url.to_string()));
body.insert(
"search".into(),
Value::String(opts.search.clone().unwrap_or_default()),
);
if let Some(limit) = opts.limit {
body.insert("limit".into(), Value::from(limit));
}
let raw = self
.send(Method::POST, path, Some(&Value::Object(body)))
.await?;
Ok(normalize_map(&raw, self.tier()))
}
pub async fn crawl(&self, params: &CrawlStartParams) -> Result<CrawlJob> {
if self.inner.api_key.is_none() {
return Err(api_key_required(
"Whole-site crawl needs an API key — set api_key or WRIT_API_KEY. \
Keyless access covers scrape and map only.",
));
}
let body = serde_json::to_value(params)
.map_err(|e| WritError::Connection(format!("serializing crawl params: {e}")))?;
let raw = self.send(Method::POST, "/api/crawl", Some(&body)).await?;
decode_crawl_job(raw)
}
pub async fn crawl_status(&self, id: i64) -> Result<CrawlJob> {
if self.inner.api_key.is_none() {
return Err(api_key_required(
"Crawl status needs an API key — set api_key or WRIT_API_KEY.",
));
}
let raw = self
.send(Method::GET, &format!("/api/crawl/{id}"), None)
.await?;
decode_crawl_job(raw)
}
pub async fn quota(&self) -> Result<Option<KeylessQuota>> {
if self.inner.api_key.is_some() {
return Ok(None);
}
let raw = self.send(Method::GET, "/v1/keyless/quota", None).await?;
Ok(Some(normalize_quota(&raw)))
}
async fn send(&self, method: Method, path: &str, json: Option<&Value>) -> Result<Value> {
let mut req = self
.inner
.http
.request(method, format!("{}{}", self.inner.base, path));
if let Some(key) = &self.inner.api_key {
req = req.header(AUTHORIZATION, format!("Bearer {key}"));
} else {
req = req.header(CLIENT_ID_HEADER, self.client_id());
}
if let Some(body) = json {
req = req.header(CONTENT_TYPE, "application/json").json(body);
}
let resp = req
.send()
.await
.map_err(|e| WritError::Connection(format!("cloud request to {path} failed: {e}")))?;
let status = resp.status().as_u16();
let text = resp
.text()
.await
.map_err(|e| WritError::Connection(format!("reading cloud response body: {e}")))?;
if !(200..300).contains(&status) {
return Err(cloud_error_from(status, &text));
}
if text.trim().is_empty() {
return Ok(Value::Object(Map::new()));
}
serde_json::from_str(&text)
.map_err(|e| WritError::Connection(format!("decoding cloud response body: {e}")))
}
fn client_id(&self) -> String {
if let Some(id) = &self.inner.client_id_override {
return id.clone();
}
self.inner
.client_id_cache
.get_or_init(load_or_mint_client_id)
.clone()
}
}
fn api_key_required(message: &str) -> WritError {
WritError::ApiKeyRequired {
status: 402,
code: "api_key_required".to_string(),
message: message.to_string(),
body: Value::Null,
}
}
fn cloud_error_from(status: u16, raw: &str) -> WritError {
let body: Value = serde_json::from_str(raw).unwrap_or_else(|_| Value::String(raw.to_string()));
let detail = body.get("detail").cloned().unwrap_or_else(|| body.clone());
let d = detail.as_object();
let field_str = |key: &str| d.and_then(|m| m.get(key)).and_then(Value::as_str);
let field_i64 = |key: &str| d.and_then(|m| m.get(key)).and_then(Value::as_i64);
let code = field_str("code")
.map(str::to_string)
.unwrap_or_else(|| code_for_status(status));
let message = field_str("message")
.map(str::to_string)
.or_else(|| detail.as_str().map(str::to_string))
.unwrap_or_else(|| format!("HTTP {status}"));
match (status, code.as_str()) {
(429, _) => WritError::RateLimited {
status,
code,
message,
reset_at: field_str("reset_at").map(str::to_string),
requests_remaining: field_i64("requests_remaining"),
pages_remaining: field_i64("pages_remaining"),
body,
},
(402, "api_key_required") => WritError::ApiKeyRequired {
status,
code,
message,
body,
},
(402, _) => WritError::InsufficientCredits {
status,
code,
message,
body,
},
_ => WritError::Api {
status,
code,
message,
body,
},
}
}
fn decode_crawl_job(raw: Value) -> Result<CrawlJob> {
serde_json::from_value(raw)
.map_err(|e| WritError::Connection(format!("decoding cloud crawl job: {e}")))
}
fn str_field(raw: &Value, key: &str) -> String {
raw.get(key)
.and_then(Value::as_str)
.unwrap_or_default()
.to_string()
}
fn opt_str_field(raw: &Value, key: &str) -> Option<String> {
raw.get(key).and_then(Value::as_str).map(str::to_string)
}
fn i64_field(raw: &Value, key: &str) -> i64 {
raw.get(key).and_then(Value::as_i64).unwrap_or(0)
}
fn normalize_quota(raw: &Value) -> KeylessQuota {
let q = raw.get("quota").unwrap_or(raw);
KeylessQuota {
tier: CloudTier::Keyless,
requests_remaining: i64_field(q, "requests_remaining"),
pages_remaining: i64_field(q, "pages_remaining"),
requests_per_day: i64_field(q, "requests_per_day"),
pages_per_day: i64_field(q, "pages_per_day"),
reset_at: str_field(q, "reset_at"),
upgrade_url: opt_str_field(q, "upgrade_url"),
}
}
fn normalize_scrape(raw: &Value, tier: CloudTier) -> ScrapeResult {
let format = {
let f = str_field(raw, "format");
if f.is_empty() {
"markdown".to_string()
} else {
f
}
};
ScrapeResult {
url: str_field(raw, "url"),
title: opt_str_field(raw, "title"),
format,
markdown: str_field(raw, "markdown"),
counts: raw
.get("counts")
.and_then(Value::as_object)
.cloned()
.unwrap_or_default(),
tier,
quota: raw.get("quota").map(|_| normalize_quota(raw)),
}
}
fn normalize_map(raw: &Value, tier: CloudTier) -> MapResult {
let urls = raw
.get("urls")
.and_then(Value::as_array)
.map(|arr| {
arr.iter()
.map(|entry| MapEntry {
url: str_field(entry, "url"),
score: entry.get("score").and_then(Value::as_f64).unwrap_or(0.0),
title: opt_str_field(entry, "title"),
})
.collect()
})
.unwrap_or_default();
let counts = raw
.get("counts")
.map(|c| MapCounts {
returned: i64_field(c, "returned"),
total: i64_field(c, "total"),
})
.unwrap_or_default();
MapResult {
url: str_field(raw, "url"),
host: opt_str_field(raw, "host"),
urls,
counts,
tier,
quota: raw.get("quota").map(|_| normalize_quota(raw)),
}
}
fn writ_home_dir() -> Option<PathBuf> {
std::env::var_os("HOME")
.or_else(|| std::env::var_os("USERPROFILE"))
.map(|home| PathBuf::from(home).join(".writ"))
}
fn load_or_mint_client_id() -> String {
let id = base64_url_nopad(&random_bytes_16());
let Some(dir) = writ_home_dir() else {
return id;
};
let file = dir.join("client_id");
if let Ok(existing) = std::fs::read_to_string(&file) {
let trimmed = existing.trim();
if !trimmed.is_empty() {
return trimmed.to_string();
}
}
let _ = std::fs::create_dir_all(&dir);
let _ = std::fs::write(&file, &id);
id
}
fn random_bytes_16() -> [u8; 16] {
use std::collections::hash_map::RandomState;
use std::hash::{BuildHasher, Hash, Hasher};
use std::time::{SystemTime, UNIX_EPOCH};
static COUNTER: AtomicU64 = AtomicU64::new(0);
let nanos = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_nanos() as u64)
.unwrap_or(0);
let seed = (
std::process::id() as u64,
nanos,
COUNTER.fetch_add(1, Ordering::Relaxed),
);
let mut out = [0u8; 16];
for (i, half) in out.chunks_mut(8).enumerate() {
let mut hasher = RandomState::new().build_hasher();
seed.hash(&mut hasher);
(i as u64).hash(&mut hasher);
half.copy_from_slice(&hasher.finish().to_le_bytes());
}
out
}
fn base64_url_nopad(bytes: &[u8]) -> String {
const ALPHABET: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_";
let mut out = String::with_capacity(bytes.len().div_ceil(3) * 4);
for chunk in bytes.chunks(3) {
let b0 = chunk[0] as u32;
let b1 = *chunk.get(1).unwrap_or(&0) as u32;
let b2 = *chunk.get(2).unwrap_or(&0) as u32;
let n = (b0 << 16) | (b1 << 8) | b2;
out.push(ALPHABET[((n >> 18) & 63) as usize] as char);
out.push(ALPHABET[((n >> 12) & 63) as usize] as char);
if chunk.len() > 1 {
out.push(ALPHABET[((n >> 6) & 63) as usize] as char);
}
if chunk.len() > 2 {
out.push(ALPHABET[(n & 63) as usize] as char);
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn tier_from_credential() {
let metered = CloudClient::builder().api_key("wt_x").build().unwrap();
assert_eq!(metered.tier(), CloudTier::Metered);
assert_eq!(metered.tier().as_str(), "metered");
let keyless = CloudClient::builder().build().unwrap();
if std::env::var_os("WRIT_API_KEY").is_none() {
assert_eq!(keyless.tier(), CloudTier::Keyless);
assert_eq!(keyless.tier().as_str(), "keyless");
}
}
#[test]
fn cloud_url_default_and_trim() {
if std::env::var_os("WRIT_CLOUD_URL").is_none() {
let c = CloudClient::builder().build().unwrap();
assert_eq!(c.base_url(), "https://api.usewrit.app");
}
let c = CloudClient::builder()
.cloud_url("https://example.test/")
.build()
.unwrap();
assert_eq!(c.base_url(), "https://example.test");
}
#[test]
fn base64_url_nopad_matches_reference() {
assert_eq!(base64_url_nopad(b""), "");
assert_eq!(base64_url_nopad(b"f"), "Zg");
assert_eq!(base64_url_nopad(b"fo"), "Zm8");
assert_eq!(base64_url_nopad(b"foo"), "Zm9v");
assert_eq!(base64_url_nopad(b"foob"), "Zm9vYg");
assert_eq!(base64_url_nopad(&[0u8; 16]).len(), 22);
}
#[test]
fn random_ids_are_distinct_and_url_safe() {
let a = base64_url_nopad(&random_bytes_16());
let b = base64_url_nopad(&random_bytes_16());
assert_ne!(a, b, "two mints must differ");
assert_eq!(a.len(), 22);
assert!(a
.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '_'));
}
#[test]
fn error_mapping_covers_each_tier_shape() {
let err = cloud_error_from(
429,
r#"{"detail":{"code":"rate_limited","message":"slow down","reset_at":"2026-07-16T00:00:00Z","requests_remaining":0,"pages_remaining":3}}"#,
);
match err {
WritError::RateLimited {
reset_at,
requests_remaining,
pages_remaining,
message,
..
} => {
assert_eq!(reset_at.as_deref(), Some("2026-07-16T00:00:00Z"));
assert_eq!(requests_remaining, Some(0));
assert_eq!(pages_remaining, Some(3));
assert_eq!(message, "slow down");
}
other => panic!("expected RateLimited, got {other:?}"),
}
let err = cloud_error_from(
402,
r#"{"detail":{"code":"api_key_required","message":"key please"}}"#,
);
assert!(
matches!(err, WritError::ApiKeyRequired { .. }),
"got {err:?}"
);
let err = cloud_error_from(
402,
r#"{"detail":{"code":"insufficient_credits","message":"broke"}}"#,
);
assert!(
matches!(err, WritError::InsufficientCredits { .. }),
"got {err:?}"
);
let err = cloud_error_from(400, r#"{"code":"bad_request","message":"nope"}"#);
match err {
WritError::Api { code, message, .. } => {
assert_eq!(code, "bad_request");
assert_eq!(message, "nope");
}
other => panic!("expected Api, got {other:?}"),
}
}
}