use std::path::PathBuf;
use std::sync::Arc;
use std::time::Duration;
use reqwest::header::{HeaderMap, HeaderValue, AUTHORIZATION};
use reqwest::{Method, Response};
use serde::de::DeserializeOwned;
use serde_json::Value;
use crate::discovery::{env_var, runtime_candidates};
use crate::error::{api_error, Result, WritError};
use crate::models::WsTicket;
use crate::resources::{
Agent, Automations, Crawl, Data, Datasets, Extractors, Files, Keys, Monitors, Personas, Runs,
Secrets, Selectors, Vault, Workflows,
};
pub(crate) const USER_AGENT: &str = concat!("writ-sdk-rust/", env!("CARGO_PKG_VERSION"));
const DEFAULT_TIMEOUT: Duration = Duration::from_secs(30);
const PROBE_TIMEOUT: Duration = Duration::from_secs(2);
pub(crate) const SSE_TIMEOUT: Duration = Duration::from_secs(24 * 60 * 60);
#[derive(Debug, Clone)]
pub struct WritAgent {
inner: Arc<Inner>,
}
#[derive(Debug)]
pub(crate) struct Inner {
pub(crate) http: reqwest::Client,
pub(crate) base_url: String,
}
#[derive(Debug, Default, Clone)]
pub struct WritAgentBuilder {
base_url: Option<String>,
token: Option<String>,
timeout: Option<Duration>,
ca_pem_file: Option<PathBuf>,
}
impl WritAgentBuilder {
pub fn base_url(mut self, url: impl Into<String>) -> Self {
self.base_url = Some(url.into());
self
}
pub fn token(mut self, token: impl Into<String>) -> Self {
self.token = Some(token.into());
self
}
pub fn timeout(mut self, timeout: Duration) -> Self {
self.timeout = Some(timeout);
self
}
pub fn ca_pem_file(mut self, path: impl Into<PathBuf>) -> Self {
self.ca_pem_file = Some(path.into());
self
}
fn http_client(&self, timeout: Duration, token: Option<&str>) -> Result<reqwest::Client> {
let mut builder = reqwest::Client::builder()
.timeout(timeout)
.user_agent(USER_AGENT);
if let Some(token) = token {
let mut headers = HeaderMap::new();
let mut auth = HeaderValue::from_str(&format!("Bearer {token}")).map_err(|_| {
WritError::Discovery("token contains characters invalid in an HTTP header".into())
})?;
auth.set_sensitive(true);
headers.insert(AUTHORIZATION, auth);
builder = builder.default_headers(headers);
}
if let Some(path) = &self.ca_pem_file {
let pem = std::fs::read(path).map_err(|e| {
WritError::Discovery(format!("cannot read ca_pem_file {}: {e}", path.display()))
})?;
let cert = reqwest::Certificate::from_pem(&pem).map_err(|e| {
WritError::Discovery(format!("invalid CA pem {}: {e}", path.display()))
})?;
builder = builder.add_root_certificate(cert);
}
builder
.build()
.map_err(|e| WritError::Discovery(format!("building http client: {e}")))
}
fn resolved(&self) -> (Option<String>, Option<String>) {
let url = self.base_url.clone().or_else(|| env_var("WRIT_API_URL"));
let token = self.token.clone().or_else(|| env_var("WRIT_TOKEN"));
(url, token)
}
fn assemble(&self, base_url: &str, token: &str) -> Result<WritAgent> {
let http = self.http_client(self.timeout.unwrap_or(DEFAULT_TIMEOUT), Some(token))?;
Ok(WritAgent {
inner: Arc::new(Inner {
http,
base_url: base_url.trim_end_matches('/').to_string(),
}),
})
}
pub fn build(self) -> Result<WritAgent> {
let (url, token) = self.resolved();
let token = token.ok_or_else(|| {
WritError::Discovery(
"no token configured — is the Writ agent running? pass .token(...) or set WRIT_TOKEN"
.into(),
)
})?;
let url = url.unwrap_or_else(|| "http://127.0.0.1:8131".to_string());
self.assemble(&url, &token)
}
pub async fn discover(self) -> Result<WritAgent> {
let (url_override, token_override) = self.resolved();
if let (Some(url), Some(token)) = (&url_override, &token_override) {
return self.assemble(url, token);
}
let candidates = runtime_candidates();
if candidates.is_empty() {
return Err(WritError::Discovery(
"no runtime.json found under $WRIT_HOME or ~/.writ — is the Writ agent running? \
pass base_url/token explicitly or set WRIT_API_URL/WRIT_TOKEN"
.into(),
));
}
let probe = self.http_client(PROBE_TIMEOUT, None)?;
let mut tried: Vec<String> = Vec::new();
for candidate in candidates {
let url = url_override
.clone()
.unwrap_or_else(|| candidate.base_url.clone());
let url = url.trim_end_matches('/').to_string();
let token = token_override
.clone()
.unwrap_or_else(|| candidate.token.clone());
let live = probe
.get(format!("{url}/v1/agent"))
.bearer_auth(&token)
.send()
.await
.map(|r| r.status().is_success())
.unwrap_or(false);
if live {
return self.assemble(&url, &token);
}
tried.push(candidate.source.display().to_string());
}
Err(WritError::Discovery(format!(
"no live Writ agent answered the probe (stale runtime.json candidates: {}) — \
is the Writ agent running? pass token=... or set WRIT_TOKEN",
tried.join(", ")
)))
}
}
impl WritAgent {
pub fn builder() -> WritAgentBuilder {
WritAgentBuilder::default()
}
pub async fn discover() -> Result<WritAgent> {
WritAgentBuilder::default().discover().await
}
pub fn base_url(&self) -> &str {
&self.inner.base_url
}
pub fn agent(&self) -> Agent<'_> {
Agent { c: &self.inner }
}
pub fn workflows(&self) -> Workflows<'_> {
Workflows { c: &self.inner }
}
pub fn runs(&self) -> Runs<'_> {
Runs { c: &self.inner }
}
pub fn monitors(&self) -> Monitors<'_> {
Monitors { c: &self.inner }
}
pub fn selectors(&self) -> Selectors<'_> {
Selectors { c: &self.inner }
}
pub fn extractors(&self) -> Extractors<'_> {
Extractors { c: &self.inner }
}
pub fn automations(&self) -> Automations<'_> {
Automations { c: &self.inner }
}
pub fn personas(&self) -> Personas<'_> {
Personas { c: &self.inner }
}
pub fn secrets(&self) -> Secrets<'_> {
Secrets { c: &self.inner }
}
pub fn vault(&self) -> Vault<'_> {
Vault { c: &self.inner }
}
pub fn files(&self) -> Files<'_> {
Files { c: &self.inner }
}
pub fn data(&self) -> Data<'_> {
Data { c: &self.inner }
}
pub fn keys(&self) -> Keys<'_> {
Keys { c: &self.inner }
}
pub fn crawl(&self) -> Crawl<'_> {
Crawl { c: &self.inner }
}
pub fn datasets(&self) -> Datasets<'_> {
Datasets { c: &self.inner }
}
pub async fn ws_ticket(&self, route: &str, channel: Option<&str>) -> Result<WsTicket> {
let mut body = serde_json::json!({ "route": route });
if let Some(channel) = channel {
body["channel"] = Value::String(channel.to_string());
}
self.inner
.send_json(Method::POST, "/v1/ws-ticket", &[], Some(&body))
.await
}
}
impl Inner {
fn url(&self, path: &str) -> String {
format!("{}{}", self.base_url, path)
}
async fn execute(&self, rb: reqwest::RequestBuilder, extra_ok: &[u16]) -> Result<Response> {
let resp = rb.send().await.map_err(WritError::from)?;
let status = resp.status();
if status.is_success() || extra_ok.contains(&status.as_u16()) {
return Ok(resp);
}
let reason = status.canonical_reason().unwrap_or("error").to_string();
let text = resp.text().await.unwrap_or_default();
Err(api_error(status.as_u16(), &reason, &text))
}
async fn decode<T: DeserializeOwned>(resp: Response) -> Result<T> {
resp.json::<T>()
.await
.map_err(|e| WritError::Connection(format!("decoding response body: {e}")))
}
pub(crate) async fn get_json<T: DeserializeOwned>(
&self,
path: &str,
query: &[(&str, &str)],
) -> Result<T> {
let rb = self.http.get(self.url(path)).query(query);
Self::decode(self.execute(rb, &[]).await?).await
}
pub(crate) async fn send_json<T: DeserializeOwned>(
&self,
method: Method,
path: &str,
query: &[(&str, &str)],
body: Option<&Value>,
) -> Result<T> {
self.send_json_allowing(method, path, query, body, &[])
.await
}
pub(crate) async fn send_json_allowing<T: DeserializeOwned>(
&self,
method: Method,
path: &str,
query: &[(&str, &str)],
body: Option<&Value>,
extra_ok: &[u16],
) -> Result<T> {
let mut rb = self.http.request(method, self.url(path)).query(query);
if let Some(body) = body {
rb = rb.json(body);
}
Self::decode(self.execute(rb, extra_ok).await?).await
}
pub(crate) async fn get_text(&self, path: &str, query: &[(&str, &str)]) -> Result<String> {
let rb = self.http.get(self.url(path)).query(query);
self.execute(rb, &[])
.await?
.text()
.await
.map_err(|e| WritError::Connection(format!("reading response body: {e}")))
}
pub(crate) async fn get_bytes(
&self,
path: &str,
query: &[(&str, &str)],
) -> Result<bytes::Bytes> {
let rb = self.http.get(self.url(path)).query(query);
self.execute(rb, &[])
.await?
.bytes()
.await
.map_err(|e| WritError::Connection(format!("reading response body: {e}")))
}
pub(crate) async fn get_stream(&self, path: &str, timeout: Duration) -> Result<Response> {
let rb = self
.http
.get(self.url(path))
.header(reqwest::header::ACCEPT, "text/event-stream")
.timeout(timeout);
self.execute(rb, &[]).await
}
pub(crate) async fn post_multipart<T: DeserializeOwned>(
&self,
path: &str,
form: reqwest::multipart::Form,
) -> Result<T> {
let rb = self.http.post(self.url(path)).multipart(form);
Self::decode(self.execute(rb, &[]).await?).await
}
}