diff options
| author | Benedikt Peetz <benedikt.peetz@b-peetz.de> | 2026-07-20 19:30:40 +0200 |
|---|---|---|
| committer | Benedikt Peetz <benedikt.peetz@b-peetz.de> | 2026-07-20 19:30:40 +0200 |
| commit | 966a80c4199a49898cc7d8641012d520ce6b2efa (patch) | |
| tree | 51029ff75842090fd1eecbea97b6f7c447e3dea9 /crates/turtle/src/client | |
| parent | chore(server): Remove warnings (diff) | |
| download | atuin-966a80c4199a49898cc7d8641012d520ce6b2efa.zip | |
chore: Commit
Diffstat (limited to 'crates/turtle/src/client')
| -rw-r--r-- | crates/turtle/src/client/mod.rs | 267 |
1 files changed, 267 insertions, 0 deletions
diff --git a/crates/turtle/src/client/mod.rs b/crates/turtle/src/client/mod.rs new file mode 100644 index 00000000..07f01e6c --- /dev/null +++ b/crates/turtle/src/client/mod.rs @@ -0,0 +1,267 @@ +use eyre::{Context as EyreContext, Result}; +use time::OffsetDateTime; +use tonic::Code; +use tonic::transport::{Channel, Endpoint, Uri}; +use tower::service_fn; + +use hyper_util::rt::TokioIo; + +#[cfg(unix)] +use tokio::net::UnixStream; + +use crate::generated::{ + self, DAEMON_PROTOCOL_VERSION, + control::{ + ForceSyncReply, ForceSyncRequest, PathsReply, PathsRequest, StatusReply, StatusRequest, + control_client::ControlClient as ControlServiceClient, + }, + history::{ + EndHistoryReply, EndHistoryRequest, HistoryEntry, HistoryRequest, StartHistoryReply, + StartHistoryRequest, TailHistoryRequest, + history_client::HistoryClient as HistoryServiceClient, + }, +}; + +pub use crate::generated::history::{HistoryEventKind, TailHistoryReply}; +use crate::history::History; + +fn normalize_optional_field(value: &str) -> Option<String> { + let trimmed = value.trim(); + if trimmed.is_empty() { + None + } else { + Some(trimmed.to_owned()) + } +} + +pub fn history_entry_to_history(entry: HistoryEntry) -> History { + let timestamp = OffsetDateTime::from_unix_timestamp_nanos(i128::from(entry.timestamp)) + .expect("Daemon history timestamp should always be valid"); + + History { + id: entry.id.into(), + timestamp, + duration: entry.duration, + exit: entry.exit, + command: entry.command, + cwd: entry.cwd, + session: entry.session, + hostname: entry.hostname, + author: entry.author, + intent: normalize_optional_field(&entry.intent), + deleted_at: None, + } +} + +#[must_use] +pub fn daemon_matches_expected(version: &str, protocol: u32) -> bool { + protocol == DAEMON_PROTOCOL_VERSION +} + +#[must_use] +pub fn daemon_mismatch_message(version: &str, protocol: u32) -> String { + if protocol == DAEMON_PROTOCOL_VERSION { + unreachable!() + } else { + format!("daemon protocol mismatch: expected {DAEMON_PROTOCOL_VERSION}, got {protocol}") + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum DaemonClientErrorKind { + Connect, + Unavailable, + Unimplemented, + Other, +} + +#[must_use] +pub fn classify_error(error: &eyre::Report) -> DaemonClientErrorKind { + for cause in error.chain() { + if cause.downcast_ref::<tonic::transport::Error>().is_some() { + return DaemonClientErrorKind::Connect; + } + + if let Some(status) = cause.downcast_ref::<tonic::Status>() { + return match status.code() { + Code::Unavailable => DaemonClientErrorKind::Unavailable, + Code::Unimplemented => DaemonClientErrorKind::Unimplemented, + _ => DaemonClientErrorKind::Other, + }; + } + } + + DaemonClientErrorKind::Other +} + +#[derive(Debug)] +pub enum Probe { + Ready(ControlClient), + NeedsRestart(String), + Unreachable(eyre::Report), +} + +/// Check if a client can reach the daemon. +pub async fn probe(path: String) -> Probe { + let mut client = match ControlClient::new(path).await { + Ok(client) => client, + Err(err) => return Probe::Unreachable(err), + }; + + match client.status().await { + Ok(status) => { + if daemon_matches_expected(&status.version, status.protocol) { + Probe::Ready(client) + } else { + Probe::NeedsRestart(daemon_mismatch_message(&status.version, status.protocol)) + } + } + Err(err) => Probe::Unreachable(err), + } +} + +// ============================================================================ +// History Client +// ============================================================================ + +#[derive(Debug)] +pub struct HistoryClient { + client: HistoryServiceClient<Channel>, +} + +pub struct Range { + pub start: OffsetDateTime, + pub end: OffsetDateTime, +} + +// Wrap the grpc client +impl HistoryClient { + #[cfg(unix)] + pub async fn new(path: String) -> Result<Self> { + use eyre::Context; + + let log_path = path.clone(); + let channel = Endpoint::try_from("http://atuin_local_daemon:0")? + .connect_with_connector(service_fn(move |_: Uri| { + let path = path.clone(); + + async move { + Ok::<_, std::io::Error>(TokioIo::new(UnixStream::connect(path.clone()).await?)) + } + })) + .await + .wrap_err_with(|| { + format!( + "failed to connect to local atuin daemon at {}. Is it running?", + &log_path + ) + })?; + + let client = HistoryServiceClient::new(channel); + + Ok(Self { client }) + } + + pub async fn start_history(&mut self, h: History) -> Result<StartHistoryReply> { + let req = StartHistoryRequest { + command: h.command, + cwd: h.cwd, + hostname: h.hostname, + session: h.session, + timestamp: h.timestamp.unix_timestamp_nanos() as u64, + author: h.author, + intent: h.intent.unwrap_or_default(), + }; + + Ok(self.client.start_history(req).await?.into_inner()) + } + + pub async fn history(&mut self, session: String, range: Option<Range>) -> Result<Vec<History>> { + let req = HistoryRequest { + session, + range: range.map(|r| generated::history::Range { + start: r.start.unix_timestamp() as u64, + end: r.end.unix_timestamp() as u64, + }), + }; + + let reply = self.client.history(req).await?.into_inner(); + + Ok(reply + .entries + .into_iter() + .map(history_entry_to_history) + .collect()) + } + + pub async fn end_history( + &mut self, + id: String, + duration: u64, + exit: i64, + ) -> Result<EndHistoryReply> { + let req = EndHistoryRequest { id, exit, duration }; + + Ok(self.client.end_history(req).await?.into_inner()) + } + + pub async fn tail_history(&mut self) -> Result<tonic::Streaming<TailHistoryReply>> { + Ok(self + .client + .tail_history(TailHistoryRequest {}) + .await? + .into_inner()) + } +} + +// ============================================================================ +// Control Client +// ============================================================================ + +/// Client for the Control gRPC service. +#[derive(Debug)] +pub struct ControlClient { + client: ControlServiceClient<Channel>, +} + +impl ControlClient { + /// Connect to the daemon's control service. + pub async fn new(path: String) -> Result<Self> { + let log_path = path.clone(); + let channel = Endpoint::try_from("http://atuin_local_daemon:0")? + .connect_with_connector(service_fn(move |_: Uri| { + let path = path.clone(); + + async move { + Ok::<_, std::io::Error>(TokioIo::new(UnixStream::connect(path.clone()).await?)) + } + })) + .await + .wrap_err_with(|| { + format!( + "failed to connect to local atuin daemon at {}. Is it running?", + &log_path + ) + })?; + + let client = ControlServiceClient::new(channel); + + Ok(Self { client }) + } + + pub async fn paths(&mut self) -> Result<PathsReply> { + Ok(self.client.paths(PathsRequest {}).await?.into_inner()) + } + + pub async fn force_sync(&mut self) -> Result<ForceSyncReply> { + Ok(self + .client + .force_sync(ForceSyncRequest {}) + .await? + .into_inner()) + } + + pub async fn status(&mut self) -> Result<StatusReply> { + Ok(self.client.status(StatusRequest {}).await?.into_inner()) + } +} |
