Skip to main content

rivetkit_client/
handle.rs

1use crate::{
2	common::{EncodingKind, RawWebSocket, TransportKind, HEADER_CONN_PARAMS, HEADER_ENCODING},
3	connection::{start_connection, ActorConnection, ActorConnectionInner},
4	protocol::{codec, query::*},
5	remote_manager::{GatewayTarget, RemoteManager},
6};
7use anyhow::{anyhow, Result};
8use bytes::Bytes;
9use reqwest::{
10	header::{HeaderMap, HeaderValue},
11	Method, Response,
12};
13use serde::Serialize;
14use serde_json::Value as JsonValue;
15use std::{
16	ops::Deref,
17	sync::{Arc, Mutex},
18	time::Duration,
19};
20
21pub use crate::protocol::codec::{QueueSendResult, QueueSendStatus};
22
23#[derive(Debug, Clone, Copy, Default)]
24pub struct SendOpts {}
25
26#[derive(Debug, Clone, Copy, Default)]
27pub struct SendAndWaitOpts {
28	pub timeout: Option<Duration>,
29}
30
31pub type QueueSendOptions = SendAndWaitOpts;
32
33pub struct ActorHandleStateless {
34	remote_manager: RemoteManager,
35	params: Option<JsonValue>,
36	encoding_kind: EncodingKind,
37	// Mutex (not RefCell) so the handle is `Sync` and `&handle` futures
38	// remain `Send` — required to call `.action(...)` from within axum
39	// middleware that needs `Send` futures.
40	query: Mutex<ActorQuery>,
41}
42
43/// An actor handle whose query has already resolved to a concrete actor ID.
44/// Raw requests through this handle never repeat Engine actor lookup.
45#[derive(Clone)]
46pub struct ResolvedActorHandle {
47	remote_manager: RemoteManager,
48	actor_id: String,
49}
50
51impl ResolvedActorHandle {
52	pub fn actor_id(&self) -> &str {
53		&self.actor_id
54	}
55
56	pub async fn fetch(
57		&self,
58		path: &str,
59		method: Method,
60		headers: HeaderMap,
61		body: Option<Bytes>,
62	) -> Result<Response> {
63		let path = normalize_fetch_path(path);
64		let target = GatewayTarget::Direct {
65			actor_id: self.actor_id.clone(),
66		};
67		self.remote_manager
68			.send_request(&target, &path, method, headers, body)
69			.await
70	}
71}
72
73impl ActorHandleStateless {
74	pub fn new(
75		remote_manager: RemoteManager,
76		params: Option<JsonValue>,
77		encoding_kind: EncodingKind,
78		query: ActorQuery,
79	) -> Self {
80		Self {
81			remote_manager,
82			params,
83			encoding_kind,
84			query: Mutex::new(query),
85		}
86	}
87
88	pub async fn action(&self, name: &str, args: Vec<JsonValue>) -> Result<JsonValue> {
89		// Resolve gateway target (query targets are resolved by the gateway).
90		let query = self.query.lock().expect("query lock poisoned").clone();
91		let target = self.remote_manager.gateway_target(&query).await?;
92
93		let body = codec::encode_http_action_request(self.encoding_kind, &args)?;
94
95		let headers = self.protocol_headers()?;
96
97		// Send request via gateway
98		let path = format!("/action/{}", urlencoding::encode(name));
99		let res = self
100			.remote_manager
101			.send_request(
102				&target,
103				&path,
104				Method::POST,
105				headers,
106				Some(Bytes::from(body)),
107			)
108			.await?;
109
110		if !res.status().is_success() {
111			let status = res.status();
112			let body = res.bytes().await?;
113			if let Ok((group, code, message, metadata)) =
114				codec::decode_http_error(self.encoding_kind, &body)
115			{
116				return Err(anyhow!(
117					"action failed ({group}/{code}): {message}, metadata={metadata:?}"
118				));
119			}
120			return Err(anyhow!("action failed: {status}"));
121		}
122
123		// Decode response
124		let output = res.bytes().await?;
125		codec::decode_http_action_response(self.encoding_kind, &output)
126	}
127
128	pub async fn send(&self, name: &str, body: impl Serialize, _opts: SendOpts) -> Result<()> {
129		self.send_queue(name, &body, false, None).await.map(|_| ())
130	}
131
132	pub async fn send_and_wait(
133		&self,
134		name: &str,
135		body: impl Serialize,
136		opts: SendAndWaitOpts,
137	) -> Result<QueueSendResult> {
138		let result = self.send_queue(name, &body, true, opts.timeout).await?;
139		result.ok_or_else(|| anyhow!("queue wait response missing"))
140	}
141
142	async fn send_queue<T: Serialize>(
143		&self,
144		name: &str,
145		body: &T,
146		wait: bool,
147		timeout: Option<Duration>,
148	) -> Result<Option<QueueSendResult>> {
149		let query = self.query.lock().expect("query lock poisoned").clone();
150		let target = self.remote_manager.gateway_target(&query).await?;
151		let timeout_ms =
152			timeout.map(|duration| u64::try_from(duration.as_millis()).unwrap_or(u64::MAX));
153		let request_body =
154			codec::encode_http_queue_request(self.encoding_kind, name, body, wait, timeout_ms)?;
155
156		let headers = self.protocol_headers()?;
157
158		let path = format!("/queue/{}", urlencoding::encode(name));
159		let res = self
160			.remote_manager
161			.send_request(
162				&target,
163				&path,
164				Method::POST,
165				headers,
166				Some(Bytes::from(request_body)),
167			)
168			.await?;
169
170		if !res.status().is_success() {
171			let status = res.status();
172			let body = res.bytes().await?;
173			if let Ok((group, code, message, metadata)) =
174				codec::decode_http_error(self.encoding_kind, &body)
175			{
176				return Err(anyhow!(
177					"queue send failed ({group}/{code}): {message}, metadata={metadata:?}"
178				));
179			}
180			return Err(anyhow!("queue send failed: {status}"));
181		}
182
183		let body = res.bytes().await?;
184		let result = codec::decode_http_queue_response(self.encoding_kind, &body)?;
185		Ok(wait.then_some(result))
186	}
187
188	pub async fn fetch(
189		&self,
190		path: &str,
191		method: Method,
192		headers: HeaderMap,
193		body: Option<Bytes>,
194	) -> Result<Response> {
195		let query = self.query.lock().expect("query lock poisoned").clone();
196		let target = self.remote_manager.gateway_target(&query).await?;
197		let path = normalize_fetch_path(path);
198		self.remote_manager
199			.send_request(&target, &path, method, headers, body)
200			.await
201	}
202
203	/// Resolves this query to a concrete actor, preserving get-or-create
204	/// behavior for creating handles.
205	pub async fn resolve_handle(&self) -> Result<ResolvedActorHandle> {
206		let query = self.query.lock().expect("query lock poisoned").clone();
207		let actor_id = self.remote_manager.resolve_actor_id(&query).await?;
208		Ok(ResolvedActorHandle {
209			remote_manager: self.remote_manager.clone(),
210			actor_id,
211		})
212	}
213
214	/// Resolves a non-creating actor query without collapsing absence into a
215	/// stringly error. Get-or-create and create queries are rejected because
216	/// their absence has different semantics.
217	pub async fn resolve_optional(&self) -> Result<Option<ResolvedActorHandle>> {
218		let query = self.query.lock().expect("query lock poisoned").clone();
219		let actor_id = match &query {
220			ActorQuery::GetForId { get_for_id } => {
221				self.remote_manager
222					.get_for_id(&get_for_id.name, &get_for_id.actor_id)
223					.await?
224			}
225			ActorQuery::GetForKey { get_for_key } => {
226				self.remote_manager
227					.get_with_key(&get_for_key.name, &get_for_key.key)
228					.await?
229			}
230			ActorQuery::GetOrCreateForKey { .. } | ActorQuery::Create { .. } => {
231				return Err(anyhow!(
232					"resolve_optional is only valid for non-creating actor queries"
233				));
234			}
235		};
236		Ok(actor_id.map(|actor_id| ResolvedActorHandle {
237			remote_manager: self.remote_manager.clone(),
238			actor_id,
239		}))
240	}
241
242	pub async fn web_socket(
243		&self,
244		path: &str,
245		protocols: Option<Vec<String>>,
246	) -> Result<RawWebSocket> {
247		let query = self.query.lock().expect("query lock poisoned").clone();
248		let target = self.remote_manager.gateway_target(&query).await?;
249		self.remote_manager
250			.open_raw_websocket(&target, path, self.params.clone(), protocols)
251			.await
252	}
253
254	pub fn gateway_url(&self) -> Result<String> {
255		let query = self.query.lock().expect("query lock poisoned").clone();
256		self.remote_manager.gateway_url(&query)
257	}
258
259	pub fn get_gateway_url(&self) -> Result<String> {
260		self.gateway_url()
261	}
262
263	pub async fn reload(&self) -> Result<()> {
264		let query = self.query.lock().expect("query lock poisoned").clone();
265		let target = self.remote_manager.gateway_target(&query).await?;
266		let res = self
267			.remote_manager
268			.send_request(
269				&target,
270				"/dynamic/reload",
271				Method::PUT,
272				HeaderMap::new(),
273				None,
274			)
275			.await?;
276		if !res.status().is_success() {
277			let status = res.status();
278			let body = res.text().await.unwrap_or_default();
279			return Err(anyhow!("reload failed with status {status}: {body}"));
280		}
281		Ok(())
282	}
283
284	pub async fn resolve(&self) -> Result<String> {
285		let query = {
286			let Ok(query) = self.query.lock() else {
287				return Err(anyhow!("Failed to lock actor query"));
288			};
289			query.clone()
290		};
291
292		match query {
293			ActorQuery::Create { .. } => Err(anyhow!("actor query cannot be create")),
294			ActorQuery::GetForId { get_for_id } => Ok(get_for_id.actor_id.clone()),
295			_ => {
296				let actor_id = self.remote_manager.resolve_actor_id(&query).await?;
297
298				// Get name from the original query
299				let name = match &query {
300					ActorQuery::GetForKey { get_for_key } => get_for_key.name.clone(),
301					ActorQuery::GetOrCreateForKey {
302						get_or_create_for_key,
303					} => get_or_create_for_key.name.clone(),
304					_ => return Err(anyhow!("unexpected query type")),
305				};
306
307				{
308					let Ok(mut query_mut) = self.query.lock() else {
309						return Err(anyhow!("Failed to lock actor query mutably"));
310					};
311
312					*query_mut = ActorQuery::GetForId {
313						get_for_id: GetForIdRequest {
314							name,
315							actor_id: actor_id.clone(),
316						},
317					};
318				}
319
320				Ok(actor_id)
321			}
322		}
323	}
324
325	fn protocol_headers(&self) -> Result<HeaderMap> {
326		let mut headers = HeaderMap::new();
327		headers.insert(
328			HEADER_ENCODING,
329			HeaderValue::from_str(self.encoding_kind.as_str())?,
330		);
331
332		if let Some(params) = &self.params {
333			headers.insert(
334				HEADER_CONN_PARAMS,
335				HeaderValue::from_str(&serde_json::to_string(params)?)?,
336			);
337		}
338
339		Ok(headers)
340	}
341}
342
343fn normalize_fetch_path(path: &str) -> String {
344	let path = path.trim_start_matches('/');
345	if path.is_empty() {
346		"/request".to_string()
347	} else {
348		format!("/request/{path}")
349	}
350}
351
352pub struct ActorHandle {
353	handle: ActorHandleStateless,
354	remote_manager: RemoteManager,
355	params: Option<JsonValue>,
356	query: ActorQuery,
357	client_shutdown_tx: Arc<tokio::sync::broadcast::Sender<()>>,
358	transport_kind: crate::TransportKind,
359	encoding_kind: EncodingKind,
360}
361
362impl ActorHandle {
363	pub fn new(
364		remote_manager: RemoteManager,
365		params: Option<JsonValue>,
366		query: ActorQuery,
367		client_shutdown_tx: Arc<tokio::sync::broadcast::Sender<()>>,
368		transport_kind: TransportKind,
369		encoding_kind: EncodingKind,
370	) -> Self {
371		let handle = ActorHandleStateless::new(
372			remote_manager.clone(),
373			params.clone(),
374			encoding_kind,
375			query.clone(),
376		);
377
378		Self {
379			handle,
380			remote_manager,
381			params,
382			query,
383			client_shutdown_tx,
384			transport_kind,
385			encoding_kind,
386		}
387	}
388
389	pub fn connect(&self) -> ActorConnection {
390		let conn = ActorConnectionInner::new(
391			self.remote_manager.clone(),
392			self.query.clone(),
393			self.transport_kind,
394			self.encoding_kind,
395			self.params.clone(),
396		);
397
398		let rx = self.client_shutdown_tx.subscribe();
399		start_connection(&conn, rx);
400
401		conn
402	}
403}
404
405impl Deref for ActorHandle {
406	type Target = ActorHandleStateless;
407
408	fn deref(&self) -> &Self::Target {
409		&self.handle
410	}
411}