diff --git a/packages/docs/content/docs/api-reference/lua/atproto-api.md b/packages/docs/content/docs/api-reference/lua/atproto-api.md index 8146401..7be0f7a 100644 --- a/packages/docs/content/docs/api-reference/lua/atproto-api.md +++ b/packages/docs/content/docs/api-reference/lua/atproto-api.md @@ -108,6 +108,104 @@ for _, uri in ipairs(uris) do end ``` +## atproto.blob_download + +```lua +local result = atproto.blob_download(did, cid) +``` + +Downloads a blob from any DID's PDS via the public `com.atproto.sync.getBlob` endpoint. No authentication is required. The blob bytes are held on the Rust side as an opaque `BlobHandle` — binary data never enters the Lua VM. + +| Parameter | Type | Description | +| --------- | ------ | ---------------------------------- | +| `did` | string | DID of the repo that owns the blob | +| `cid` | string | CID of the blob to download | + +**Returns:** A table with: + +| Field | Type | Description | +| ---------- | ---------- | -------------------------------------------------------- | +| `handle` | BlobHandle | Opaque handle to the blob bytes (pass to `blob_upload`) | +| `mimeType` | string | Content type from the PDS response (e.g. `"image/png"`) | +| `size` | number | Size of the blob in bytes | + +If the content-type header is missing from the PDS response, `mimeType` defaults to `"application/octet-stream"`. + +**Throws** on any non-2xx response from the PDS, including 404 (blob not found) and 429 (rate limited). Retry logic is the script's responsibility. + +**Availability:** All script contexts (queries, procedures, record scripts). + +### BlobHandle methods + +The `BlobHandle` userdata exposes two methods: + +| Method | Returns | Description | +| ------------- | ------- | --------------------------------- | +| `:size()` | number | Size of the blob in bytes | +| `:mime_type()` | string | MIME type of the blob | + +### Examples + +```lua +-- Download a blob and inspect it +local result = atproto.blob_download("did:plc:abc123", "bafyreie...") +log("downloaded " .. result.size .. " bytes, type: " .. result.mimeType) + +-- The handle can also be queried directly +log("handle size: " .. result.handle:size()) +log("handle mime: " .. result.handle:mime_type()) +``` + +## atproto.blob_upload + +```lua +local response = atproto.blob_upload(handle, content_type) +``` + +Uploads blob bytes to the caller's PDS via authenticated `com.atproto.repo.uploadBlob`. The `handle` must be a `BlobHandle` from `blob_download`. + +| Parameter | Type | Description | +| -------------- | ---------- | -------------------------------------------- | +| `handle` | BlobHandle | Opaque blob handle from `blob_download` | +| `content_type` | string | MIME type for the upload (e.g. `"image/png"`) | + +**Returns:** The PDS `uploadBlob` response, which contains a `blob` field with the new blob reference: + +```lua +{ + blob = { + ["$type"] = "blob", + ref = { ["$link"] = "" }, + mimeType = "image/png", + size = 12345 + } +} +``` + +**Throws** on any error, including 429 (rate limited) and authentication failures. Retry logic is the script's responsibility. + +**Availability:** Procedure scripts only. Returns `nil` in query and record script contexts (no PDS auth available). + +### Examples + +```lua +-- Copy a blob from one repo to another +local downloaded = atproto.blob_download(source_did, old_cid) +local uploaded = atproto.blob_upload(downloaded.handle, downloaded.mimeType) + +-- Use the new blob ref in a record +local new_cid = uploaded.blob.ref["$link"] + +-- Migrate all blobs in a media array +for _, item in ipairs(record.media) do + if item.blob and item.blob.ref then + local dl = atproto.blob_download(source_did, item.blob.ref["$link"]) + local ul = atproto.blob_upload(dl.handle, dl.mimeType) + item.blob = ul.blob + end +end +``` + ## atproto.sign ```lua diff --git a/src/lua/atproto_api.rs b/src/lua/atproto_api.rs index aeb047c..be944ba 100644 --- a/src/lua/atproto_api.rs +++ b/src/lua/atproto_api.rs @@ -1,3 +1,4 @@ +use axum::body::Bytes; use mlua::{Lua, LuaSerdeExt, Result as LuaResult}; use std::sync::Arc; @@ -5,6 +6,22 @@ use crate::AppState; use crate::db::{adapt_sql, now_rfc3339}; use crate::profile; +/// Opaque handle to blob bytes stored on the Rust side. +/// Lua scripts receive this from `atproto.blob_download()` and pass it +/// to `atproto.blob_upload()` — the binary data never enters the Lua VM. +#[derive(Clone)] +pub(crate) struct BlobHandle { + pub data: Bytes, + pub mime_type: String, +} + +impl mlua::UserData for BlobHandle { + fn add_methods>(methods: &mut M) { + methods.add_method("size", |_, this, ()| Ok(this.data.len())); + methods.add_method("mime_type", |_, this, ()| Ok(this.mime_type.clone())); + } +} + /// Register the `atproto` table with AT Protocol utility functions. /// /// When `caller_did` is provided, the `atproto.sign(record)` function is @@ -32,6 +49,68 @@ pub fn register_atproto_api( atproto_table.set("resolve_service_endpoint", resolve_fn)?; + // atproto.blob_download(did, cid) -> { handle = BlobHandle, mimeType = string, size = number } + { + let state_clone = state.clone(); + let blob_download_fn = + lua.create_async_function(move |lua, (did, cid): (String, String)| { + let state = state_clone.clone(); + async move { + let pds_endpoint = + profile::resolve_pds_endpoint(&state.http, &state.config.plc_url, &did) + .await + .map_err(|e| { + mlua::Error::runtime(format!( + "blob_download: failed to resolve PDS for {did}: {e}" + )) + })?; + + let url = format!( + "{}/xrpc/com.atproto.sync.getBlob?did={}&cid={}", + pds_endpoint, + urlencoding::encode(&did), + urlencoding::encode(&cid), + ); + + let response = state.http.get(&url).send().await.map_err(|e| { + mlua::Error::runtime(format!("blob_download: request failed: {e}")) + })?; + + let status = response.status(); + if !status.is_success() { + return Err(mlua::Error::runtime(format!( + "blob_download: PDS returned {status} for did={did} cid={cid}" + ))); + } + + let mime_type = response + .headers() + .get("content-type") + .and_then(|v| v.to_str().ok()) + .unwrap_or("application/octet-stream") + .to_string(); + + let bytes = response.bytes().await.map_err(|e| { + mlua::Error::runtime(format!("blob_download: failed to read body: {e}")) + })?; + + let size = bytes.len(); + let handle = BlobHandle { + data: bytes, + mime_type: mime_type.clone(), + }; + + let result = lua.create_table()?; + result.set("handle", lua.create_userdata(handle)?)?; + result.set("mimeType", mime_type)?; + result.set("size", size)?; + + Ok(mlua::Value::Table(result)) + } + })?; + atproto_table.set("blob_download", blob_download_fn)?; + } + // get_labels(uri) -> array of { src, uri, val, cts } let state_clone = state.clone(); let get_labels_fn = lua.create_async_function(move |lua, uri: String| { @@ -472,6 +551,123 @@ pub fn register_atproto_api( Ok(()) } +/// Register blob upload capability on the existing `atproto` table. +/// +/// Called only in procedure execution contexts where PDS auth is +/// available. `blob_download` is registered in `register_atproto_api` +/// (available everywhere); `blob_upload` needs caller credentials to +/// write to the caller's PDS. +pub(crate) fn register_atproto_blob_api( + lua: &Lua, + state: Arc, + claims: Arc, + pds_auth: Arc, +) -> LuaResult<()> { + let atproto_table: mlua::Table = lua.globals().get("atproto")?; + + let upload_fn = lua.create_async_function( + move |lua, (handle, content_type): (mlua::AnyUserData, String)| { + let state = state.clone(); + let claims = claims.clone(); + let pds_auth = pds_auth.clone(); + async move { + let blob_handle = handle.borrow::().map_err(|_| { + mlua::Error::runtime( + "blob_upload: first argument must be a BlobHandle from blob_download()", + ) + })?; + let blob_bytes = blob_handle.data.clone(); + drop(blob_handle); + + let result = + upload_blob_to_pds(&state, claims.did(), &pds_auth, &content_type, blob_bytes) + .await + .map_err(|e| mlua::Error::runtime(format!("blob_upload: {e}")))?; + + lua.to_value(&result) + } + }, + )?; + atproto_table.set("blob_upload", upload_fn)?; + + Ok(()) +} + +async fn upload_blob_to_pds( + state: &AppState, + caller_did: &str, + pds_auth: &crate::repo::PdsAuth, + content_type: &str, + blob_bytes: Bytes, +) -> Result { + use crate::error::AppError; + use crate::repo::PdsAuth; + + match pds_auth { + PdsAuth::OAuth(session) => { + use atrium_xrpc::{ + InputDataOrBytes, OutputDataOrBytes, XrpcClient, XrpcRequest, http::Method, + }; + + let request = XrpcRequest { + method: Method::POST, + nsid: "com.atproto.repo.uploadBlob".to_string(), + parameters: None::<()>, + input: Some(InputDataOrBytes::<()>::Bytes(blob_bytes.to_vec())), + encoding: Some(content_type.to_string()), + }; + + let result: Result< + OutputDataOrBytes, + atrium_xrpc::Error, + > = session.send_xrpc(&request).await; + + match result { + Ok(OutputDataOrBytes::Data(data)) => Ok(data), + Ok(OutputDataOrBytes::Bytes(bytes)) => serde_json::from_slice(&bytes) + .map_err(|e| AppError::Internal(format!("invalid uploadBlob response: {e}"))), + Err(e) => Err(AppError::Internal(format!("PDS uploadBlob failed: {e}"))), + } + } + PdsAuth::Dpop { + api_client_id, + dpop_key_id, + encryption_key, + } => { + let resp = crate::oauth::pds_write::dpop_pds_post_blob( + &state.http, + &state.db, + state.db_backend, + encryption_key, + &state.oauth, + &state.config.plc_url, + api_client_id, + caller_did, + dpop_key_id, + content_type, + blob_bytes, + ) + .await?; + + let status = resp.status(); + let body = resp + .bytes() + .await + .map_err(|e| AppError::Internal(format!("failed to read upload response: {e}")))?; + + if !status.is_success() { + let body_str = String::from_utf8_lossy(&body); + return Err(AppError::Internal(format!( + "PDS uploadBlob returned {status}: {body_str}" + ))); + } + + serde_json::from_slice(&body) + .map_err(|e| AppError::Internal(format!("invalid uploadBlob response: {e}"))) + } + } +} + #[cfg(test)] mod tests { use super::*; @@ -749,4 +945,230 @@ mod tests { let result: bool = lua.load(chunk).eval_async().await.unwrap(); assert!(result); } + + #[tokio::test] + async fn blob_handle_exposes_size_and_mime() { + let state = test_state_with_plc(""); + let lua = mlua::Lua::new(); + register_atproto_api(&lua, Arc::new(state), None).unwrap(); + + let handle = BlobHandle { + data: axum::body::Bytes::from_static(b"hello world"), + mime_type: "text/plain".to_string(), + }; + lua.globals() + .set("test_handle", lua.create_userdata(handle).unwrap()) + .unwrap(); + + let size: usize = lua + .load("return test_handle:size()") + .eval_async() + .await + .unwrap(); + assert_eq!(size, 11); + + let mime: String = lua + .load("return test_handle:mime_type()") + .eval_async() + .await + .unwrap(); + assert_eq!(mime, "text/plain"); + } + + #[tokio::test] + async fn blob_download_returns_handle_and_metadata() { + let mock = wiremock::MockServer::start().await; + + let did_doc = serde_json::json!({ + "id": "did:plc:blobsource", + "service": [{ + "id": "#atproto_pds", + "type": "AtprotoPersonalDataServer", + "serviceEndpoint": mock.uri() + }] + }); + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/did:plc:blobsource")) + .respond_with(wiremock::ResponseTemplate::new(200).set_body_json(&did_doc)) + .mount(&mock) + .await; + + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/xrpc/com.atproto.sync.getBlob")) + .and(wiremock::matchers::query_param("did", "did:plc:blobsource")) + .and(wiremock::matchers::query_param("cid", "bafytest123")) + .respond_with( + wiremock::ResponseTemplate::new(200) + .insert_header("content-type", "image/png") + .set_body_bytes(vec![0x89, 0x50, 0x4E, 0x47]), + ) + .mount(&mock) + .await; + + let state = test_state_with_plc(&mock.uri()); + let lua = mlua::Lua::new(); + register_atproto_api(&lua, Arc::new(state), None).unwrap(); + + let chunk = r#" + local result = atproto.blob_download("did:plc:blobsource", "bafytest123") + return { + size = result.handle:size(), + mimeType = result.mimeType, + } + "#; + let result: mlua::Table = lua.load(chunk).eval_async().await.unwrap(); + assert_eq!(result.get::("size").unwrap(), 4); + assert_eq!(result.get::("mimeType").unwrap(), "image/png"); + } + + #[tokio::test] + async fn blob_download_throws_on_404() { + let mock = wiremock::MockServer::start().await; + + let did_doc = serde_json::json!({ + "id": "did:plc:blobsource", + "service": [{ + "id": "#atproto_pds", + "type": "AtprotoPersonalDataServer", + "serviceEndpoint": mock.uri() + }] + }); + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/did:plc:blobsource")) + .respond_with(wiremock::ResponseTemplate::new(200).set_body_json(&did_doc)) + .mount(&mock) + .await; + + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/xrpc/com.atproto.sync.getBlob")) + .respond_with(wiremock::ResponseTemplate::new(404)) + .mount(&mock) + .await; + + let state = test_state_with_plc(&mock.uri()); + let lua = mlua::Lua::new(); + register_atproto_api(&lua, Arc::new(state), None).unwrap(); + + let result: Result = lua + .load(r#"return atproto.blob_download("did:plc:blobsource", "bafymissing")"#) + .eval_async() + .await; + assert!(result.is_err()); + } + + #[tokio::test] + async fn blob_download_defaults_mime_to_octet_stream() { + let mock = wiremock::MockServer::start().await; + + let did_doc = serde_json::json!({ + "id": "did:plc:blobsource", + "service": [{ + "id": "#atproto_pds", + "type": "AtprotoPersonalDataServer", + "serviceEndpoint": mock.uri() + }] + }); + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/did:plc:blobsource")) + .respond_with(wiremock::ResponseTemplate::new(200).set_body_json(&did_doc)) + .mount(&mock) + .await; + + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/xrpc/com.atproto.sync.getBlob")) + .respond_with(wiremock::ResponseTemplate::new(200).set_body_bytes(vec![0xFF, 0xD8])) + .mount(&mock) + .await; + + let state = test_state_with_plc(&mock.uri()); + let lua = mlua::Lua::new(); + register_atproto_api(&lua, Arc::new(state), None).unwrap(); + + let chunk = r#" + local result = atproto.blob_download("did:plc:blobsource", "bafynoheader") + return result.mimeType + "#; + let result: String = lua.load(chunk).eval_async().await.unwrap(); + assert_eq!(result, "application/octet-stream"); + } + + #[tokio::test] + async fn blob_download_throws_on_429() { + let mock = wiremock::MockServer::start().await; + + let did_doc = serde_json::json!({ + "id": "did:plc:blobsource", + "service": [{ + "id": "#atproto_pds", + "type": "AtprotoPersonalDataServer", + "serviceEndpoint": mock.uri() + }] + }); + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/did:plc:blobsource")) + .respond_with(wiremock::ResponseTemplate::new(200).set_body_json(&did_doc)) + .mount(&mock) + .await; + + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/xrpc/com.atproto.sync.getBlob")) + .respond_with(wiremock::ResponseTemplate::new(429)) + .mount(&mock) + .await; + + let state = test_state_with_plc(&mock.uri()); + let lua = mlua::Lua::new(); + register_atproto_api(&lua, Arc::new(state), None).unwrap(); + + let result: Result = lua + .load(r#"return atproto.blob_download("did:plc:blobsource", "bafyratelimit")"#) + .eval_async() + .await; + assert!(result.is_err()); + let err_msg = result.unwrap_err().to_string(); + assert!( + err_msg.contains("429"), + "error should mention 429 status: {err_msg}" + ); + } + + #[tokio::test] + async fn blob_download_throws_on_did_resolution_failure() { + let mock = wiremock::MockServer::start().await; + + wiremock::Mock::given(wiremock::matchers::method("GET")) + .and(wiremock::matchers::path("/did:plc:nonexistent")) + .respond_with(wiremock::ResponseTemplate::new(404)) + .mount(&mock) + .await; + + let state = test_state_with_plc(&mock.uri()); + let lua = mlua::Lua::new(); + register_atproto_api(&lua, Arc::new(state), None).unwrap(); + + let result: Result = lua + .load(r#"return atproto.blob_download("did:plc:nonexistent", "bafytest")"#) + .eval_async() + .await; + assert!(result.is_err()); + let err_msg = result.unwrap_err().to_string(); + assert!( + err_msg.contains("resolve PDS"), + "error should mention PDS resolution failure: {err_msg}" + ); + } + + #[tokio::test] + async fn blob_upload_throws_without_auth() { + let state = test_state_with_plc(""); + let lua = mlua::Lua::new(); + register_atproto_api(&lua, Arc::new(state), None).unwrap(); + + let has_upload: bool = lua + .load("return atproto.blob_upload ~= nil") + .eval_async() + .await + .unwrap(); + assert!(!has_upload); + } } diff --git a/src/lua/execute.rs b/src/lua/execute.rs index 7ee6f3c..8a77b62 100644 --- a/src/lua/execute.rs +++ b/src/lua/execute.rs @@ -259,6 +259,35 @@ pub async fn execute_procedure_script( return Err(AppError::Internal(error_message)); } + if let Err(e) = atproto_api::register_atproto_blob_api( + &lua, + state_arc.clone(), + claims_arc.clone(), + pds_auth_arc.clone(), + ) { + let error_message = format!("failed to register atproto blob API: {e}"); + log_event( + &state.db, + EventLog { + event_type: "script.error".to_string(), + severity: Severity::Error, + actor_did: Some(claims.did().to_string()), + subject: Some(method.to_string()), + detail: serde_json::json!({ + "error": error_message, + "script_source": script_source, + "input": input_json, + "caller_did": claims.did(), + "method": method, + "duration_ms": start.elapsed().as_millis() as u64, + }), + }, + backend, + ) + .await; + return Err(AppError::Internal(error_message)); + } + if let Err(e) = record::register_record_api( &lua, state_arc.clone(),