diff --git a/src/lua/context.rs b/src/lua/context.rs index 161c58d..282e46d 100644 --- a/src/lua/context.rs +++ b/src/lua/context.rs @@ -24,11 +24,16 @@ pub fn set_query_context( method: &str, params: &HashMap, collection: &str, + caller_did: Option<&str>, ) -> LuaResult<()> { let globals = lua.globals(); globals.set("method", method.to_string())?; globals.set("params", lua.to_value(params)?)?; globals.set("collection", collection.to_string())?; + match caller_did { + Some(did) => globals.set("caller_did", did.to_string())?, + None => globals.set("caller_did", mlua::Value::Nil)?, + } Ok(()) } @@ -131,7 +136,14 @@ mod tests { let mut params = HashMap::new(); params.insert("limit".to_string(), json!("10")); params.insert("cursor".to_string(), json!("abc")); - set_query_context(&lua, "com.example.listThings", ¶ms, "com.example.thing").unwrap(); + set_query_context( + &lua, + "com.example.listThings", + ¶ms, + "com.example.thing", + Some("did:plc:test"), + ) + .unwrap(); let globals = lua.globals(); assert_eq!( diff --git a/src/lua/execute.rs b/src/lua/execute.rs index 558eb24..daa7469 100644 --- a/src/lua/execute.rs +++ b/src/lua/execute.rs @@ -341,6 +341,7 @@ pub async fn execute_query_script( params: &HashMap, lexicon: &ParsedLexicon, script: &str, + claims: Option<&Claims>, ) -> Result { let start = Instant::now(); let collection = lexicon.target_collection.as_deref().unwrap_or_default(); @@ -413,7 +414,9 @@ pub async fn execute_query_script( return Err(AppError::Internal(error_message)); } - if let Err(e) = context::set_query_context(&lua, method, params, collection) { + if let Err(e) = + context::set_query_context(&lua, method, params, collection, claims.map(|c| c.did())) + { let error_message = format!("failed to set context: {e}"); log_event( &state.db, diff --git a/src/xrpc/mod.rs b/src/xrpc/mod.rs index 785dbf9..e4ad56d 100644 --- a/src/xrpc/mod.rs +++ b/src/xrpc/mod.rs @@ -3,8 +3,9 @@ mod query; use axum::Json; use axum::body::Body; -use axum::extract::{Path, RawQuery, State}; +use axum::extract::{FromRequestParts, Path, RawQuery, State}; use axum::http::StatusCode; +use axum::http::request::Parts; use axum::response::Response; use serde_json::Value; use std::collections::HashMap; @@ -109,9 +110,11 @@ pub async fn xrpc_get( State(state): State, Path(method): Path, RawQuery(raw_query): RawQuery, + mut parts: Parts, ) -> Result { let raw_query = raw_query.unwrap_or_default(); let params = parse_query_params(&raw_query); + let claims = Claims::from_request_parts(&mut parts, &state).await.ok(); let lexicon = match state.lexicons.get(&method).await { Some(l) => l, @@ -126,7 +129,7 @@ pub async fn xrpc_get( ))); } - query::handle_query(&state, &method, ¶ms, &lexicon).await + query::handle_query(&state, &method, ¶ms, &lexicon, claims.as_ref()).await } /// Catch-all POST handler for XRPC procedures. diff --git a/src/xrpc/query.rs b/src/xrpc/query.rs index 70b3561..a25bb3e 100644 --- a/src/xrpc/query.rs +++ b/src/xrpc/query.rs @@ -4,6 +4,7 @@ use serde_json::{Value, json}; use std::collections::HashMap; use crate::AppState; +use crate::auth::Claims; use crate::error::AppError; pub(super) async fn handle_query( @@ -11,9 +12,11 @@ pub(super) async fn handle_query( method: &str, params: &HashMap, lexicon: &crate::lexicon::ParsedLexicon, + claims: Option<&Claims>, ) -> Result { if let Some(ref script) = lexicon.script { - return crate::lua::execute_query_script(state, method, params, lexicon, script).await; + return crate::lua::execute_query_script(state, method, params, lexicon, script, claims) + .await; } // Single-record query: has a `uri` parameter