Skip to content

Commit e355fba

Browse files
authored
refactor: Update AI backend to support multiple conversation types (#361)
1 parent a428f68 commit e355fba

22 files changed

Lines changed: 1147 additions & 681 deletions
Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,2 @@
1+
ALTER TABLE ai_sessions
2+
DROP COLUMN version;
Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,2 @@
1+
ALTER TABLE ai_sessions
2+
ADD COLUMN version INTEGER NOT NULL DEFAULT 0;

backend/src/ai/client.rs

Lines changed: 115 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,115 @@
1+
use std::{future::Future, ops::Deref, pin::Pin, sync::Arc};
2+
3+
use genai::{
4+
adapter::AdapterKind,
5+
resolver::{AuthData, Endpoint, ServiceTargetResolver},
6+
ClientConfig, ModelIden, ServiceTarget,
7+
};
8+
9+
use crate::secret_cache::{SecretCache, SecretCacheError};
10+
11+
#[derive(Debug, thiserror::Error)]
12+
pub enum AtuinAIClientError {
13+
#[error("Failed to get credential: {0}")]
14+
CredentialError(#[from] SecretCacheError),
15+
}
16+
17+
/// A wrapper around a genai::Client that includes Atuin's custom service target resolver
18+
pub struct AtuinAIClient {
19+
client: genai::Client,
20+
#[allow(dead_code)] // this will be used to fetch provider API keys in the future
21+
secret_cache: Arc<SecretCache>,
22+
}
23+
24+
impl AtuinAIClient {
25+
pub fn new(secret_cache: Arc<SecretCache>) -> Self {
26+
let secret_cache_clone = secret_cache.clone();
27+
let target_resolver =
28+
ServiceTargetResolver::from_resolver_async_fn(move |service_target| {
29+
resolve_service_target(service_target, secret_cache_clone.clone())
30+
});
31+
let client = genai::Client::builder()
32+
.with_config(ClientConfig::default().with_service_target_resolver(target_resolver))
33+
.build();
34+
35+
Self {
36+
client,
37+
secret_cache,
38+
}
39+
}
40+
}
41+
42+
impl Deref for AtuinAIClient {
43+
type Target = genai::Client;
44+
45+
fn deref(&self) -> &genai::Client {
46+
&self.client
47+
}
48+
}
49+
50+
fn resolve_service_target(
51+
mut service_target: ServiceTarget,
52+
secret_cache: Arc<SecretCache>,
53+
) -> Pin<Box<dyn Future<Output = Result<ServiceTarget, genai::resolver::Error>> + Send>> {
54+
Box::pin(async move {
55+
let model_name = service_target.model.model_name.to_string();
56+
let parts = model_name.splitn(3, "::").collect::<Vec<&str>>();
57+
58+
if parts.len() != 3 {
59+
return Err(genai::resolver::Error::Custom(format!(
60+
"Invalid Atuin Desktop model identifier format: {}",
61+
model_name
62+
)));
63+
}
64+
65+
// Set the adapter kind based on the provider
66+
let adapter_kind = match parts[0] {
67+
"atuinhub" => AdapterKind::Anthropic,
68+
"claude" => AdapterKind::Anthropic,
69+
"openai" => AdapterKind::OpenAI,
70+
"ollama" => AdapterKind::Ollama,
71+
_ => {
72+
return Err(genai::resolver::Error::Custom(format!(
73+
"Invalid provider identifier: {}",
74+
parts[0]
75+
)))
76+
}
77+
};
78+
79+
// Set the API key, if any, for the provider
80+
let key = get_api_key(&secret_cache, adapter_kind, parts[0] == "atuinhub")
81+
.await
82+
.map_err(|e| genai::resolver::Error::Custom(e.to_string()))?;
83+
84+
if let Some(key) = key {
85+
let auth = AuthData::Key(key);
86+
service_target.auth = auth;
87+
} else {
88+
service_target.auth = AuthData::Key("".to_string());
89+
}
90+
91+
// Set the specific model
92+
let model_id = ModelIden::new(adapter_kind, parts[1]);
93+
service_target.model = model_id;
94+
95+
// Set the endpoint, if any, for the model
96+
if parts[2] != "default" {
97+
service_target.endpoint = Endpoint::from_owned(parts[2].to_string());
98+
}
99+
100+
Ok(service_target)
101+
})
102+
}
103+
104+
async fn get_api_key(
105+
_secret_cache: &SecretCache,
106+
_adapter_kind: AdapterKind,
107+
is_hub: bool,
108+
) -> Result<Option<String>, AtuinAIClientError> {
109+
// todo
110+
if is_hub {
111+
return Ok(None);
112+
}
113+
114+
Ok(None)
115+
}

backend/src/ai/fsm.rs

Lines changed: 0 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -11,8 +11,6 @@ use genai::chat::{ChatMessage, ContentPart, MessageContent, ToolCall, ToolRespon
1111
use serde::{Deserialize, Serialize};
1212
use ts_rs::TS;
1313

14-
use crate::ai::{session::ChargeTarget, types::ModelSelection};
15-
1614
// ============================================================================
1715
// FSM State
1816
// ============================================================================
@@ -84,15 +82,6 @@ pub struct Context {
8482
/// Events that drive state transitions.
8583
#[derive(Debug, Clone)]
8684
pub enum Event {
87-
/// Change the charge target.
88-
ChargeTargetChange(ChargeTarget),
89-
90-
/// Change active user.
91-
UserChange(String),
92-
93-
/// User requests a new model.
94-
ModelChange(ModelSelection),
95-
9685
/// User submitted a message.
9786
UserMessage(ChatMessage),
9887

@@ -124,15 +113,6 @@ pub enum Event {
124113
/// The FSM returns these; it never executes them directly.
125114
#[derive(Debug, Clone)]
126115
pub enum Effect {
127-
/// Change the model.
128-
ModelChange(ModelSelection),
129-
130-
/// Change the charge target.
131-
ChargeTargetChange(ChargeTarget),
132-
133-
/// Change active user.
134-
UserChange(String),
135-
136116
/// Start a new request to the model.
137117
/// Caller should use context.conversation to build the actual request.
138118
StartRequest,
@@ -279,18 +259,6 @@ impl Agent {
279259
pub fn handle(&mut self, event: Event) -> Transition {
280260
// State-specific transitions
281261
match (&self.state, event) {
282-
// ================================================================
283-
// Model change
284-
// ================================================================
285-
// TODO: should we handle this in any state, or require Idle?
286-
(_, Event::ModelChange(model)) => Transition::single(Effect::ModelChange(model)),
287-
288-
(_, Event::ChargeTargetChange(charge_target)) => {
289-
Transition::single(Effect::ChargeTargetChange(charge_target))
290-
}
291-
292-
(_, Event::UserChange(user)) => Transition::single(Effect::UserChange(user)),
293-
294262
// ================================================================
295263
// User messages
296264
// ================================================================

0 commit comments

Comments
 (0)