feat: AI based permission checks
This commit is contained in:
@@ -5,6 +5,7 @@ use std::fs;
|
||||
use std::path::Path;
|
||||
use time::OffsetDateTime;
|
||||
|
||||
use crate::config::PermissionMode;
|
||||
use crate::schema::{a2ui_messages, messages, projects, sessions};
|
||||
|
||||
pub const MIGRATIONS: EmbeddedMigrations = embed_migrations!("migrations");
|
||||
@@ -42,6 +43,8 @@ pub struct Session {
|
||||
state: String,
|
||||
pub compacted_summary: Option<String>,
|
||||
last_used: i64,
|
||||
/// Raw column value; read it through [`Session::permission_mode`].
|
||||
permission_mode: String,
|
||||
}
|
||||
|
||||
impl Session {
|
||||
@@ -49,6 +52,10 @@ impl Session {
|
||||
SessionState::from_id(&self.state).unwrap_or_default()
|
||||
}
|
||||
|
||||
pub fn permission_mode(&self) -> PermissionMode {
|
||||
PermissionMode::from_id(&self.permission_mode).unwrap_or_default()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub fn fixture(id: i32, project_id: i32, title: &str, state: SessionState) -> Self {
|
||||
Self {
|
||||
@@ -61,6 +68,7 @@ impl Session {
|
||||
state: state.as_id().to_owned(),
|
||||
compacted_summary: None,
|
||||
last_used: 0,
|
||||
permission_mode: PermissionMode::default().as_id().to_owned(),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -109,6 +117,7 @@ struct NewSession<'a> {
|
||||
project_id: i32,
|
||||
title: &'a str,
|
||||
last_used: i64,
|
||||
permission_mode: &'a str,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Identifiable, Queryable, Selectable)]
|
||||
@@ -330,18 +339,36 @@ impl Database {
|
||||
.map_err(|error: diesel::result::Error| error.to_string())
|
||||
}
|
||||
|
||||
pub fn create_session(&mut self, project_id: i32, title: &str) -> Result<Session, String> {
|
||||
pub fn create_session(
|
||||
&mut self,
|
||||
project_id: i32,
|
||||
title: &str,
|
||||
permission_mode: PermissionMode,
|
||||
) -> Result<Session, String> {
|
||||
diesel::insert_into(sessions::table)
|
||||
.values(NewSession {
|
||||
project_id,
|
||||
title,
|
||||
last_used: OffsetDateTime::now_utc().unix_timestamp(),
|
||||
permission_mode: permission_mode.as_id(),
|
||||
})
|
||||
.returning(Session::as_returning())
|
||||
.get_result(&mut self.connection)
|
||||
.map_err(|error| error.to_string())
|
||||
}
|
||||
|
||||
pub fn set_session_permission_mode(
|
||||
&mut self,
|
||||
session_id: i32,
|
||||
permission_mode: PermissionMode,
|
||||
) -> Result<(), String> {
|
||||
diesel::update(sessions::table.find(session_id))
|
||||
.set(sessions::permission_mode.eq(permission_mode.as_id()))
|
||||
.execute(&mut self.connection)
|
||||
.map(|_| ())
|
||||
.map_err(|error| error.to_string())
|
||||
}
|
||||
|
||||
pub fn touch_session(&mut self, session_id: i32) -> Result<(), String> {
|
||||
touch_session(&mut self.connection, session_id).map_err(|error| error.to_string())
|
||||
}
|
||||
@@ -685,19 +712,31 @@ mod tests {
|
||||
|
||||
let project = database.create_project("DS4", "/tmp/ds4").unwrap();
|
||||
let first = database
|
||||
.create_session(project.id, "First session")
|
||||
.create_session(project.id, "First session", PermissionMode::Heuristic)
|
||||
.unwrap();
|
||||
database
|
||||
.create_session(project.id, "Second session")
|
||||
let second = database
|
||||
.create_session(project.id, "Second session", PermissionMode::Ai)
|
||||
.unwrap();
|
||||
database.delete_session(first.id).unwrap();
|
||||
let loaded = database.load_projects().unwrap();
|
||||
assert_eq!(loaded[0].sessions[0].title, "Second session");
|
||||
assert_eq!(loaded[0].sessions[0].state(), SessionState::Normal);
|
||||
assert_eq!(loaded[0].sessions[0].permission_mode(), PermissionMode::Ai);
|
||||
database
|
||||
.set_session_permission_mode(second.id, PermissionMode::Heuristic)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
database.load_projects().unwrap()[0].sessions[0].permission_mode(),
|
||||
PermissionMode::Heuristic
|
||||
);
|
||||
|
||||
let ordinary = loaded[0].sessions[0].id;
|
||||
let pinned = database.create_session(project.id, "Pinned").unwrap();
|
||||
let archived = database.create_session(project.id, "Archived").unwrap();
|
||||
let pinned = database
|
||||
.create_session(project.id, "Pinned", PermissionMode::Heuristic)
|
||||
.unwrap();
|
||||
let archived = database
|
||||
.create_session(project.id, "Archived", PermissionMode::Heuristic)
|
||||
.unwrap();
|
||||
database
|
||||
.set_session_state(pinned.id, SessionState::Pinned)
|
||||
.unwrap();
|
||||
@@ -752,9 +791,15 @@ mod tests {
|
||||
let project = database
|
||||
.create_project("DS4", "/tmp/ds4-session-order")
|
||||
.unwrap();
|
||||
let older = database.create_session(project.id, "Older").unwrap();
|
||||
let newer = database.create_session(project.id, "Newer").unwrap();
|
||||
let pinned = database.create_session(project.id, "Pinned").unwrap();
|
||||
let older = database
|
||||
.create_session(project.id, "Older", PermissionMode::Heuristic)
|
||||
.unwrap();
|
||||
let newer = database
|
||||
.create_session(project.id, "Newer", PermissionMode::Heuristic)
|
||||
.unwrap();
|
||||
let pinned = database
|
||||
.create_session(project.id, "Pinned", PermissionMode::Heuristic)
|
||||
.unwrap();
|
||||
database
|
||||
.set_session_state(pinned.id, SessionState::Pinned)
|
||||
.unwrap();
|
||||
@@ -792,7 +837,9 @@ mod tests {
|
||||
let project = database
|
||||
.create_project("DS4", "/tmp/ds4-reactivate")
|
||||
.unwrap();
|
||||
let session = database.create_session(project.id, "Archived").unwrap();
|
||||
let session = database
|
||||
.create_session(project.id, "Archived", PermissionMode::Heuristic)
|
||||
.unwrap();
|
||||
database
|
||||
.set_session_state(session.id, SessionState::Archived)
|
||||
.unwrap();
|
||||
@@ -830,7 +877,9 @@ mod tests {
|
||||
let project = database
|
||||
.create_project("DS4", "/tmp/ds4-a2ui-dismiss")
|
||||
.unwrap();
|
||||
let session = database.create_session(project.id, "A2UI").unwrap();
|
||||
let session = database
|
||||
.create_session(project.id, "A2UI", PermissionMode::Heuristic)
|
||||
.unwrap();
|
||||
let first = database
|
||||
.start_chat_turn(session.id, "First", None, &[], false)
|
||||
.unwrap()
|
||||
@@ -908,7 +957,9 @@ mod tests {
|
||||
let path = std::env::temp_dir().join(format!("ds4-chat-{id}.sqlite3"));
|
||||
let mut database = Database::open(&path).unwrap();
|
||||
let project = database.create_project("DS4", "/tmp/ds4-chat").unwrap();
|
||||
let session = database.create_session(project.id, "Chat").unwrap();
|
||||
let session = database
|
||||
.create_session(project.id, "Chat", PermissionMode::Heuristic)
|
||||
.unwrap();
|
||||
let mut opening = database
|
||||
.start_chat_turn(
|
||||
session.id,
|
||||
|
||||
Reference in New Issue
Block a user