mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-26 23:50:17 +00:00
Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
47f80ddc22 | ||
|
|
d2098f1030 | ||
|
|
79a2c5d9dd |
@@ -54,7 +54,7 @@ jobs:
|
|||||||
- group: features
|
- group: features
|
||||||
files: "tests/e2e/scenarios/test_skills.py tests/e2e/scenarios/test_tool_approval.py tests/e2e/scenarios/test_webhook.py"
|
files: "tests/e2e/scenarios/test_skills.py tests/e2e/scenarios/test_tool_approval.py tests/e2e/scenarios/test_webhook.py"
|
||||||
- group: extensions
|
- group: extensions
|
||||||
files: "tests/e2e/scenarios/test_extensions.py tests/e2e/scenarios/test_extension_oauth.py tests/e2e/scenarios/test_oauth_url_parameters.py tests/e2e/scenarios/test_telegram_token_validation.py tests/e2e/scenarios/test_telegram_hot_activation.py tests/e2e/scenarios/test_wasm_lifecycle.py tests/e2e/scenarios/test_tool_execution.py tests/e2e/scenarios/test_pairing.py tests/e2e/scenarios/test_mcp_auth_flow.py tests/e2e/scenarios/test_oauth_credential_fallback.py tests/e2e/scenarios/test_routine_oauth_credential_injection.py"
|
files: "tests/e2e/scenarios/test_extensions.py tests/e2e/scenarios/test_extension_oauth.py tests/e2e/scenarios/test_telegram_token_validation.py tests/e2e/scenarios/test_telegram_hot_activation.py tests/e2e/scenarios/test_wasm_lifecycle.py tests/e2e/scenarios/test_tool_execution.py tests/e2e/scenarios/test_pairing.py tests/e2e/scenarios/test_mcp_auth_flow.py tests/e2e/scenarios/test_oauth_credential_fallback.py tests/e2e/scenarios/test_routine_oauth_credential_injection.py"
|
||||||
- group: routines
|
- group: routines
|
||||||
files: "tests/e2e/scenarios/test_owner_scope.py tests/e2e/scenarios/test_routine_event_batch.py"
|
files: "tests/e2e/scenarios/test_owner_scope.py tests/e2e/scenarios/test_routine_event_batch.py"
|
||||||
steps:
|
steps:
|
||||||
|
|||||||
Generated
+12
-12
@@ -157,7 +157,7 @@ version = "1.1.5"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc"
|
checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"windows-sys 0.60.2",
|
"windows-sys 0.61.2",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -168,7 +168,7 @@ checksum = "291e6a250ff86cd4a820112fb8898808a366d8f9f58ce16d1f538353ad55747d"
|
|||||||
dependencies = [
|
dependencies = [
|
||||||
"anstyle",
|
"anstyle",
|
||||||
"once_cell_polyfill",
|
"once_cell_polyfill",
|
||||||
"windows-sys 0.60.2",
|
"windows-sys 0.61.2",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -2136,7 +2136,7 @@ dependencies = [
|
|||||||
"libc",
|
"libc",
|
||||||
"option-ext",
|
"option-ext",
|
||||||
"redox_users 0.5.2",
|
"redox_users 0.5.2",
|
||||||
"windows-sys 0.59.0",
|
"windows-sys 0.61.2",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -2323,7 +2323,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||||||
checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb"
|
checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"libc",
|
"libc",
|
||||||
"windows-sys 0.52.0",
|
"windows-sys 0.61.2",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -4134,7 +4134,7 @@ version = "0.50.3"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5"
|
checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"windows-sys 0.59.0",
|
"windows-sys 0.61.2",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -5472,7 +5472,7 @@ dependencies = [
|
|||||||
"errno",
|
"errno",
|
||||||
"libc",
|
"libc",
|
||||||
"linux-raw-sys 0.12.1",
|
"linux-raw-sys 0.12.1",
|
||||||
"windows-sys 0.52.0",
|
"windows-sys 0.61.2",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -6154,7 +6154,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||||||
checksum = "3a766e1110788c36f4fa1c2b71b387a7815aa65f88ce0229841826633d93723e"
|
checksum = "3a766e1110788c36f4fa1c2b71b387a7815aa65f88ce0229841826633d93723e"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"libc",
|
"libc",
|
||||||
"windows-sys 0.60.2",
|
"windows-sys 0.61.2",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -6354,9 +6354,9 @@ checksum = "55937e1799185b12863d447f42597ed69d9928686b8d88a1df17376a097d8369"
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "tar"
|
name = "tar"
|
||||||
version = "0.4.45"
|
version = "0.4.44"
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "22692a6476a21fa75fdfc11d452fda482af402c008cdbaf3476414e122040973"
|
checksum = "1d863878d212c87a19c1a610eb53bb01fe12951c0501cf5a0d65f724914a667a"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"filetime",
|
"filetime",
|
||||||
"libc",
|
"libc",
|
||||||
@@ -6379,7 +6379,7 @@ dependencies = [
|
|||||||
"getrandom 0.4.2",
|
"getrandom 0.4.2",
|
||||||
"once_cell",
|
"once_cell",
|
||||||
"rustix 1.1.4",
|
"rustix 1.1.4",
|
||||||
"windows-sys 0.52.0",
|
"windows-sys 0.61.2",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -7179,7 +7179,7 @@ checksum = "f2f6fb2847f6742cd76af783a2a2c49e9375d0a111c7bef6f71cd9e738c72d6e"
|
|||||||
dependencies = [
|
dependencies = [
|
||||||
"memoffset",
|
"memoffset",
|
||||||
"tempfile",
|
"tempfile",
|
||||||
"windows-sys 0.60.2",
|
"windows-sys 0.61.2",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -8029,7 +8029,7 @@ version = "0.1.11"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22"
|
checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"windows-sys 0.48.0",
|
"windows-sys 0.61.2",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
|
|||||||
+1
-1
@@ -161,7 +161,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
|||||||
| `config` | ✅ | ✅ | - | Read/write config plus validate/path helpers |
|
| `config` | ✅ | ✅ | - | Read/write config plus validate/path helpers |
|
||||||
| `backup` | ✅ | ❌ | P3 | Create/verify local backup archives |
|
| `backup` | ✅ | ❌ | P3 | Create/verify local backup archives |
|
||||||
| `channels` | ✅ | 🚧 | P2 | `list` implemented; `enable`/`disable`/`status` deferred pending config source unification |
|
| `channels` | ✅ | 🚧 | P2 | `list` implemented; `enable`/`disable`/`status` deferred pending config source unification |
|
||||||
| `models` | ✅ | 🚧 | P1 | `models list [<provider>]` (`--verbose`, `--json`; fetches live model list when provider specified), `models status` (`--json`), `models set <model>`, `models set-provider <provider> [--model model]` (alias normalization, config.toml + .env persistence). Remaining: `set` doesn't validate model against live list. |
|
| `models` | ✅ | 🚧 | - | Model selector in TUI |
|
||||||
| `status` | ✅ | ✅ | - | System status (enriched session details) |
|
| `status` | ✅ | ✅ | - | System status (enriched session details) |
|
||||||
| `agents` | ✅ | ❌ | P3 | Multi-agent management |
|
| `agents` | ✅ | ❌ | P3 | Multi-agent management |
|
||||||
| `sessions` | ✅ | ❌ | P3 | Session listing (shows subagent models) |
|
| `sessions` | ✅ | ❌ | P3 | Session listing (shows subagent models) |
|
||||||
|
|||||||
+12
-22
@@ -44,7 +44,7 @@ pub struct JobMonitorRoute {
|
|||||||
/// the main agent's context window).
|
/// the main agent's context window).
|
||||||
pub fn spawn_job_monitor(
|
pub fn spawn_job_monitor(
|
||||||
job_id: Uuid,
|
job_id: Uuid,
|
||||||
event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>,
|
event_rx: broadcast::Receiver<(Uuid, SseEvent)>,
|
||||||
inject_tx: mpsc::Sender<IncomingMessage>,
|
inject_tx: mpsc::Sender<IncomingMessage>,
|
||||||
route: JobMonitorRoute,
|
route: JobMonitorRoute,
|
||||||
) -> JoinHandle<()> {
|
) -> JoinHandle<()> {
|
||||||
@@ -56,7 +56,7 @@ pub fn spawn_job_monitor(
|
|||||||
/// jobs don't stay `InProgress` forever in the `ContextManager`.
|
/// jobs don't stay `InProgress` forever in the `ContextManager`.
|
||||||
pub fn spawn_job_monitor_with_context(
|
pub fn spawn_job_monitor_with_context(
|
||||||
job_id: Uuid,
|
job_id: Uuid,
|
||||||
mut event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>,
|
mut event_rx: broadcast::Receiver<(Uuid, SseEvent)>,
|
||||||
inject_tx: mpsc::Sender<IncomingMessage>,
|
inject_tx: mpsc::Sender<IncomingMessage>,
|
||||||
route: JobMonitorRoute,
|
route: JobMonitorRoute,
|
||||||
context_manager: Option<Arc<ContextManager>>,
|
context_manager: Option<Arc<ContextManager>>,
|
||||||
@@ -68,7 +68,7 @@ pub fn spawn_job_monitor_with_context(
|
|||||||
|
|
||||||
loop {
|
loop {
|
||||||
match event_rx.recv().await {
|
match event_rx.recv().await {
|
||||||
Ok((ev_job_id, _user_id, event)) => {
|
Ok((ev_job_id, event)) => {
|
||||||
if ev_job_id != job_id {
|
if ev_job_id != job_id {
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
@@ -162,7 +162,7 @@ pub fn spawn_job_monitor_with_context(
|
|||||||
/// inject messages into) but we still need to free the `max_jobs` slot.
|
/// inject messages into) but we still need to free the `max_jobs` slot.
|
||||||
pub fn spawn_completion_watcher(
|
pub fn spawn_completion_watcher(
|
||||||
job_id: Uuid,
|
job_id: Uuid,
|
||||||
mut event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>,
|
mut event_rx: broadcast::Receiver<(Uuid, SseEvent)>,
|
||||||
context_manager: Arc<ContextManager>,
|
context_manager: Arc<ContextManager>,
|
||||||
) -> JoinHandle<()> {
|
) -> JoinHandle<()> {
|
||||||
let short_id = job_id.to_string()[..8].to_string();
|
let short_id = job_id.to_string()[..8].to_string();
|
||||||
@@ -170,9 +170,7 @@ pub fn spawn_completion_watcher(
|
|||||||
tokio::spawn(async move {
|
tokio::spawn(async move {
|
||||||
loop {
|
loop {
|
||||||
match event_rx.recv().await {
|
match event_rx.recv().await {
|
||||||
Ok((ev_job_id, _user_id, SseEvent::JobResult { status, .. }))
|
Ok((ev_job_id, SseEvent::JobResult { status, .. })) if ev_job_id == job_id => {
|
||||||
if ev_job_id == job_id =>
|
|
||||||
{
|
|
||||||
let target = if status == "completed" {
|
let target = if status == "completed" {
|
||||||
JobState::Completed
|
JobState::Completed
|
||||||
} else {
|
} else {
|
||||||
@@ -229,7 +227,7 @@ mod tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_monitor_forwards_assistant_messages() {
|
async fn test_monitor_forwards_assistant_messages() {
|
||||||
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
|
let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
|
||||||
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
|
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
|
||||||
|
|
||||||
let job_id = Uuid::new_v4();
|
let job_id = Uuid::new_v4();
|
||||||
@@ -239,7 +237,6 @@ mod tests {
|
|||||||
event_tx
|
event_tx
|
||||||
.send((
|
.send((
|
||||||
job_id,
|
job_id,
|
||||||
"test-user".to_string(),
|
|
||||||
SseEvent::JobMessage {
|
SseEvent::JobMessage {
|
||||||
job_id: job_id.to_string(),
|
job_id: job_id.to_string(),
|
||||||
role: "assistant".to_string(),
|
role: "assistant".to_string(),
|
||||||
@@ -262,7 +259,7 @@ mod tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_monitor_ignores_other_jobs() {
|
async fn test_monitor_ignores_other_jobs() {
|
||||||
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
|
let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
|
||||||
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
|
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
|
||||||
|
|
||||||
let job_id = Uuid::new_v4();
|
let job_id = Uuid::new_v4();
|
||||||
@@ -273,7 +270,6 @@ mod tests {
|
|||||||
event_tx
|
event_tx
|
||||||
.send((
|
.send((
|
||||||
other_job_id,
|
other_job_id,
|
||||||
"test-user".to_string(),
|
|
||||||
SseEvent::JobMessage {
|
SseEvent::JobMessage {
|
||||||
job_id: other_job_id.to_string(),
|
job_id: other_job_id.to_string(),
|
||||||
role: "assistant".to_string(),
|
role: "assistant".to_string(),
|
||||||
@@ -293,7 +289,7 @@ mod tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_monitor_exits_on_job_result() {
|
async fn test_monitor_exits_on_job_result() {
|
||||||
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
|
let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
|
||||||
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
|
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
|
||||||
|
|
||||||
let job_id = Uuid::new_v4();
|
let job_id = Uuid::new_v4();
|
||||||
@@ -303,7 +299,6 @@ mod tests {
|
|||||||
event_tx
|
event_tx
|
||||||
.send((
|
.send((
|
||||||
job_id,
|
job_id,
|
||||||
"test-user".to_string(),
|
|
||||||
SseEvent::JobResult {
|
SseEvent::JobResult {
|
||||||
job_id: job_id.to_string(),
|
job_id: job_id.to_string(),
|
||||||
status: "completed".to_string(),
|
status: "completed".to_string(),
|
||||||
@@ -329,7 +324,7 @@ mod tests {
|
|||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_monitor_skips_tool_events() {
|
async fn test_monitor_skips_tool_events() {
|
||||||
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
|
let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
|
||||||
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
|
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
|
||||||
|
|
||||||
let job_id = Uuid::new_v4();
|
let job_id = Uuid::new_v4();
|
||||||
@@ -339,7 +334,6 @@ mod tests {
|
|||||||
event_tx
|
event_tx
|
||||||
.send((
|
.send((
|
||||||
job_id,
|
job_id,
|
||||||
"test-user".to_string(),
|
|
||||||
SseEvent::JobToolUse {
|
SseEvent::JobToolUse {
|
||||||
job_id: job_id.to_string(),
|
job_id: job_id.to_string(),
|
||||||
tool_name: "shell".to_string(),
|
tool_name: "shell".to_string(),
|
||||||
@@ -352,7 +346,6 @@ mod tests {
|
|||||||
event_tx
|
event_tx
|
||||||
.send((
|
.send((
|
||||||
job_id,
|
job_id,
|
||||||
"test-user".to_string(),
|
|
||||||
SseEvent::JobMessage {
|
SseEvent::JobMessage {
|
||||||
job_id: job_id.to_string(),
|
job_id: job_id.to_string(),
|
||||||
role: "user".to_string(),
|
role: "user".to_string(),
|
||||||
@@ -409,7 +402,7 @@ mod tests {
|
|||||||
.await
|
.await
|
||||||
.unwrap();
|
.unwrap();
|
||||||
|
|
||||||
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
|
let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
|
||||||
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
|
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
|
||||||
|
|
||||||
let handle = spawn_job_monitor_with_context(
|
let handle = spawn_job_monitor_with_context(
|
||||||
@@ -424,7 +417,6 @@ mod tests {
|
|||||||
event_tx
|
event_tx
|
||||||
.send((
|
.send((
|
||||||
job_id,
|
job_id,
|
||||||
"test-user".to_string(),
|
|
||||||
SseEvent::JobResult {
|
SseEvent::JobResult {
|
||||||
job_id: job_id.to_string(),
|
job_id: job_id.to_string(),
|
||||||
status: "completed".to_string(),
|
status: "completed".to_string(),
|
||||||
@@ -458,7 +450,7 @@ mod tests {
|
|||||||
.await
|
.await
|
||||||
.unwrap();
|
.unwrap();
|
||||||
|
|
||||||
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
|
let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
|
||||||
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
|
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
|
||||||
|
|
||||||
let handle = spawn_job_monitor_with_context(
|
let handle = spawn_job_monitor_with_context(
|
||||||
@@ -473,7 +465,6 @@ mod tests {
|
|||||||
event_tx
|
event_tx
|
||||||
.send((
|
.send((
|
||||||
job_id,
|
job_id,
|
||||||
"test-user".to_string(),
|
|
||||||
SseEvent::JobResult {
|
SseEvent::JobResult {
|
||||||
job_id: job_id.to_string(),
|
job_id: job_id.to_string(),
|
||||||
status: "failed".to_string(),
|
status: "failed".to_string(),
|
||||||
@@ -507,13 +498,12 @@ mod tests {
|
|||||||
.await
|
.await
|
||||||
.unwrap();
|
.unwrap();
|
||||||
|
|
||||||
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
|
let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
|
||||||
let handle = spawn_completion_watcher(job_id, event_tx.subscribe(), Arc::clone(&cm));
|
let handle = spawn_completion_watcher(job_id, event_tx.subscribe(), Arc::clone(&cm));
|
||||||
|
|
||||||
event_tx
|
event_tx
|
||||||
.send((
|
.send((
|
||||||
job_id,
|
job_id,
|
||||||
"test-user".to_string(),
|
|
||||||
SseEvent::JobResult {
|
SseEvent::JobResult {
|
||||||
job_id: job_id.to_string(),
|
job_id: job_id.to_string(),
|
||||||
status: "completed".to_string(),
|
status: "completed".to_string(),
|
||||||
|
|||||||
+154
-32
@@ -992,8 +992,16 @@ impl Agent {
|
|||||||
{
|
{
|
||||||
// Put it back and return error
|
// Put it back and return error
|
||||||
let mut sess = session.lock().await;
|
let mut sess = session.lock().await;
|
||||||
if let Some(thread) = sess.threads.get_mut(&thread_id) {
|
match sess.threads.get_mut(&thread_id) {
|
||||||
thread.await_approval(pending);
|
Some(thread) => {
|
||||||
|
thread.await_approval(pending);
|
||||||
|
}
|
||||||
|
None => {
|
||||||
|
tracing::warn!(
|
||||||
|
%thread_id,
|
||||||
|
"Thread disappeared while restoring pending approval after request ID mismatch"
|
||||||
|
);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
return Ok(SubmissionResult::error(
|
return Ok(SubmissionResult::error(
|
||||||
"Request ID mismatch. Use the correct request ID.",
|
"Request ID mismatch. Use the correct request ID.",
|
||||||
@@ -1015,8 +1023,19 @@ impl Agent {
|
|||||||
// Reset thread state to processing
|
// Reset thread state to processing
|
||||||
{
|
{
|
||||||
let mut sess = session.lock().await;
|
let mut sess = session.lock().await;
|
||||||
if let Some(thread) = sess.threads.get_mut(&thread_id) {
|
match sess.threads.get_mut(&thread_id) {
|
||||||
thread.state = ThreadState::Processing;
|
Some(thread) => {
|
||||||
|
thread.state = ThreadState::Processing;
|
||||||
|
}
|
||||||
|
None => {
|
||||||
|
tracing::error!(
|
||||||
|
%thread_id,
|
||||||
|
"Thread disappeared while setting state to Processing during approval"
|
||||||
|
);
|
||||||
|
return Ok(SubmissionResult::error(
|
||||||
|
"Internal error: thread no longer exists",
|
||||||
|
));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1100,13 +1119,21 @@ impl Agent {
|
|||||||
// Record sanitized result in thread
|
// Record sanitized result in thread
|
||||||
{
|
{
|
||||||
let mut sess = session.lock().await;
|
let mut sess = session.lock().await;
|
||||||
if let Some(thread) = sess.threads.get_mut(&thread_id)
|
match sess.threads.get_mut(&thread_id) {
|
||||||
&& let Some(turn) = thread.last_turn_mut()
|
Some(thread) => {
|
||||||
{
|
if let Some(turn) = thread.last_turn_mut() {
|
||||||
if is_tool_error {
|
if is_tool_error {
|
||||||
turn.record_tool_error(result_content.clone());
|
turn.record_tool_error(result_content.clone());
|
||||||
} else {
|
} else {
|
||||||
turn.record_tool_result(serde_json::json!(result_content));
|
turn.record_tool_result(serde_json::json!(result_content));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
None => {
|
||||||
|
tracing::error!(
|
||||||
|
%thread_id,
|
||||||
|
"Thread disappeared while recording tool result during approval"
|
||||||
|
);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -1354,13 +1381,22 @@ impl Agent {
|
|||||||
// Record sanitized result in thread
|
// Record sanitized result in thread
|
||||||
{
|
{
|
||||||
let mut sess = session.lock().await;
|
let mut sess = session.lock().await;
|
||||||
if let Some(thread) = sess.threads.get_mut(&thread_id)
|
match sess.threads.get_mut(&thread_id) {
|
||||||
&& let Some(turn) = thread.last_turn_mut()
|
Some(thread) => {
|
||||||
{
|
if let Some(turn) = thread.last_turn_mut() {
|
||||||
if is_deferred_error {
|
if is_deferred_error {
|
||||||
turn.record_tool_error(deferred_content.clone());
|
turn.record_tool_error(deferred_content.clone());
|
||||||
} else {
|
} else {
|
||||||
turn.record_tool_result(serde_json::json!(deferred_content));
|
turn.record_tool_result(serde_json::json!(deferred_content));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
None => {
|
||||||
|
tracing::error!(
|
||||||
|
%thread_id,
|
||||||
|
tool_name = %tc.name,
|
||||||
|
"Thread disappeared while recording deferred tool result during approval"
|
||||||
|
);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -1413,8 +1449,19 @@ impl Agent {
|
|||||||
|
|
||||||
{
|
{
|
||||||
let mut sess = session.lock().await;
|
let mut sess = session.lock().await;
|
||||||
if let Some(thread) = sess.threads.get_mut(&thread_id) {
|
match sess.threads.get_mut(&thread_id) {
|
||||||
thread.await_approval(new_pending);
|
Some(thread) => {
|
||||||
|
thread.await_approval(new_pending);
|
||||||
|
}
|
||||||
|
None => {
|
||||||
|
tracing::error!(
|
||||||
|
%thread_id,
|
||||||
|
"Thread disappeared while setting up deferred tool approval"
|
||||||
|
);
|
||||||
|
return Ok(SubmissionResult::error(
|
||||||
|
"Internal error: thread no longer exists",
|
||||||
|
));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1546,17 +1593,28 @@ impl Agent {
|
|||||||
);
|
);
|
||||||
{
|
{
|
||||||
let mut sess = session.lock().await;
|
let mut sess = session.lock().await;
|
||||||
if let Some(thread) = sess.threads.get_mut(&thread_id) {
|
match sess.threads.get_mut(&thread_id) {
|
||||||
thread.clear_pending_approval();
|
Some(thread) => {
|
||||||
thread.complete_turn(&rejection);
|
thread.clear_pending_approval();
|
||||||
// User message already persisted at turn start; save rejection response
|
thread.complete_turn(&rejection);
|
||||||
self.persist_assistant_response(
|
// User message already persisted at turn start; save rejection response
|
||||||
thread_id,
|
self.persist_assistant_response(
|
||||||
&message.channel,
|
thread_id,
|
||||||
&message.user_id,
|
&message.channel,
|
||||||
&rejection,
|
&message.user_id,
|
||||||
)
|
&rejection,
|
||||||
.await;
|
)
|
||||||
|
.await;
|
||||||
|
}
|
||||||
|
None => {
|
||||||
|
tracing::error!(
|
||||||
|
%thread_id,
|
||||||
|
"Thread disappeared during approval rejection"
|
||||||
|
);
|
||||||
|
return Ok(SubmissionResult::error(
|
||||||
|
"Internal error: thread no longer exists",
|
||||||
|
));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1646,7 +1704,7 @@ impl Agent {
|
|||||||
};
|
};
|
||||||
|
|
||||||
match ext_mgr
|
match ext_mgr
|
||||||
.configure_token(&pending.extension_name, token, &message.user_id)
|
.configure_token(&pending.extension_name, token)
|
||||||
.await
|
.await
|
||||||
{
|
{
|
||||||
Ok(result) if result.activated => {
|
Ok(result) if result.activated => {
|
||||||
@@ -2098,6 +2156,70 @@ mod tests {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_approval_on_missing_thread_should_error() {
|
||||||
|
// Regression for #1487: when a thread disappears from the session
|
||||||
|
// during approval processing, the code must return a visible error
|
||||||
|
// rather than silently succeeding.
|
||||||
|
//
|
||||||
|
// We can't call process_approval() directly (requires full Agent),
|
||||||
|
// so we simulate the exact code pattern used in the rejection and
|
||||||
|
// state-setting paths: lock session, match on get_mut, verify the
|
||||||
|
// None arm produces an error.
|
||||||
|
use crate::agent::session::{Session, Thread, ThreadState};
|
||||||
|
use std::sync::Arc;
|
||||||
|
use tokio::sync::Mutex;
|
||||||
|
use uuid::Uuid;
|
||||||
|
|
||||||
|
let thread_id = Uuid::new_v4();
|
||||||
|
let session_id = Uuid::new_v4();
|
||||||
|
let session = Arc::new(Mutex::new(Session::new("test-user")));
|
||||||
|
|
||||||
|
// Scenario 1: Thread never existed
|
||||||
|
{
|
||||||
|
let sess = session.lock().await;
|
||||||
|
let result = match sess.threads.get(&thread_id) {
|
||||||
|
Some(_) => Ok("processed"),
|
||||||
|
None => Err("Internal error: thread no longer exists"),
|
||||||
|
};
|
||||||
|
assert!(result.is_err());
|
||||||
|
assert_eq!(
|
||||||
|
result.unwrap_err(),
|
||||||
|
"Internal error: thread no longer exists"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Scenario 2: Thread existed then was removed (simulates disappearance
|
||||||
|
// between lock acquisitions -- the TOCTOU window this fix addresses)
|
||||||
|
{
|
||||||
|
let mut sess = session.lock().await;
|
||||||
|
let mut thread = Thread::with_id(thread_id, session_id);
|
||||||
|
thread.start_turn("pending approval");
|
||||||
|
thread.state = ThreadState::AwaitingApproval;
|
||||||
|
sess.threads.insert(thread_id, thread);
|
||||||
|
}
|
||||||
|
{
|
||||||
|
let mut sess = session.lock().await;
|
||||||
|
// Simulate thread disappearing (e.g., pruned by another task)
|
||||||
|
sess.threads.remove(&thread_id);
|
||||||
|
|
||||||
|
// The rejection path must detect this and return an error
|
||||||
|
let result = match sess.threads.get_mut(&thread_id) {
|
||||||
|
Some(thread) => {
|
||||||
|
thread.clear_pending_approval();
|
||||||
|
thread.complete_turn("rejected");
|
||||||
|
Ok("rejection persisted")
|
||||||
|
}
|
||||||
|
None => Err("Internal error: thread no longer exists"),
|
||||||
|
};
|
||||||
|
assert!(result.is_err());
|
||||||
|
assert_eq!(
|
||||||
|
result.unwrap_err(),
|
||||||
|
"Internal error: thread no longer exists"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_queue_cap_rejects_at_capacity() {
|
fn test_queue_cap_rejects_at_capacity() {
|
||||||
use crate::agent::session::{MAX_PENDING_MESSAGES, Thread, ThreadState};
|
use crate::agent::session::{MAX_PENDING_MESSAGES, Thread, ThreadState};
|
||||||
|
|||||||
+2
-31
@@ -327,7 +327,7 @@ impl AppBuilder {
|
|||||||
.with_search_config(&self.config.search);
|
.with_search_config(&self.config.search);
|
||||||
|
|
||||||
if let Some(ref emb) = embeddings {
|
if let Some(ref emb) = embeddings {
|
||||||
ws = ws.with_embeddings_cached(emb.clone(), emb_cache_config.clone());
|
ws = ws.with_embeddings_cached(emb.clone(), emb_cache_config);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Wire workspace-level settings (read scopes, memory layers)
|
// Wire workspace-level settings (read scopes, memory layers)
|
||||||
@@ -341,36 +341,7 @@ impl AppBuilder {
|
|||||||
}
|
}
|
||||||
ws = ws.with_memory_layers(self.config.workspace.memory_layers.clone());
|
ws = ws.with_memory_layers(self.config.workspace.memory_layers.clone());
|
||||||
let ws = Arc::new(ws);
|
let ws = Arc::new(ws);
|
||||||
|
tools.register_memory_tools(Arc::clone(&ws));
|
||||||
// Detect multi-tenant mode: when GATEWAY_USER_TOKENS is configured,
|
|
||||||
// each authenticated user needs their own workspace scope. Use
|
|
||||||
// PerUserWorkspaceResolver to create per-user workspaces on demand
|
|
||||||
// instead of sharing the startup workspace across all users.
|
|
||||||
let is_multi_tenant = self
|
|
||||||
.config
|
|
||||||
.channels
|
|
||||||
.gateway
|
|
||||||
.as_ref()
|
|
||||||
.is_some_and(|gw| gw.user_tokens.is_some());
|
|
||||||
|
|
||||||
if is_multi_tenant {
|
|
||||||
let resolver = Arc::new(
|
|
||||||
crate::tools::builtin::memory::PerUserWorkspaceResolver::new(
|
|
||||||
Arc::clone(db),
|
|
||||||
embeddings.clone(),
|
|
||||||
emb_cache_config,
|
|
||||||
self.config.search.clone(),
|
|
||||||
self.config.workspace.clone(),
|
|
||||||
),
|
|
||||||
);
|
|
||||||
tools.register_memory_tools_with_resolver(resolver);
|
|
||||||
tracing::info!(
|
|
||||||
"Memory tools configured with per-user workspace resolver (multi-tenant mode)"
|
|
||||||
);
|
|
||||||
} else {
|
|
||||||
tools.register_memory_tools(Arc::clone(&ws));
|
|
||||||
}
|
|
||||||
|
|
||||||
Some(ws)
|
Some(ws)
|
||||||
} else {
|
} else {
|
||||||
None
|
None
|
||||||
|
|||||||
+22
-383
@@ -1,133 +1,17 @@
|
|||||||
//! Bearer token authentication middleware for the web gateway.
|
//! Bearer token authentication middleware for the web gateway.
|
||||||
//!
|
|
||||||
//! Supports multi-user mode: each token maps to a `UserIdentity` that carries
|
|
||||||
//! the user_id. The identity is inserted into request extensions so downstream
|
|
||||||
//! handlers can extract it via `AuthenticatedUser`.
|
|
||||||
|
|
||||||
use std::collections::HashMap;
|
|
||||||
|
|
||||||
use axum::{
|
use axum::{
|
||||||
extract::{FromRequestParts, Request, State},
|
extract::{Request, State},
|
||||||
http::{HeaderMap, Method, StatusCode, request::Parts},
|
http::{HeaderMap, Method, StatusCode},
|
||||||
middleware::Next,
|
middleware::Next,
|
||||||
response::{IntoResponse, Response},
|
response::{IntoResponse, Response},
|
||||||
};
|
};
|
||||||
use sha2::{Digest, Sha256};
|
|
||||||
use subtle::ConstantTimeEq;
|
use subtle::ConstantTimeEq;
|
||||||
|
|
||||||
/// Identity resolved from a bearer token.
|
/// Shared auth state injected via axum middleware state.
|
||||||
#[derive(Debug, Clone)]
|
|
||||||
pub struct UserIdentity {
|
|
||||||
pub user_id: String,
|
|
||||||
/// Additional user scopes this identity can read from.
|
|
||||||
pub workspace_read_scopes: Vec<String>,
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Hash a token with SHA-256 for constant-size, timing-safe storage.
|
|
||||||
fn hash_token(token: &str) -> [u8; 32] {
|
|
||||||
let mut hasher = Sha256::new();
|
|
||||||
hasher.update(token.as_bytes());
|
|
||||||
hasher.finalize().into()
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Multi-user auth state: maps token hashes to user identities.
|
|
||||||
///
|
|
||||||
/// Tokens are SHA-256 hashed on construction so they are never stored in
|
|
||||||
/// plaintext. Authentication compares fixed-size (32-byte) digests using
|
|
||||||
/// constant-time comparison, eliminating both length-oracle timing leaks
|
|
||||||
/// and accidental token exposure in memory dumps.
|
|
||||||
///
|
|
||||||
/// In single-user mode (the default), contains exactly one entry.
|
|
||||||
#[derive(Clone)]
|
#[derive(Clone)]
|
||||||
pub struct MultiAuthState {
|
pub struct AuthState {
|
||||||
/// Maps SHA-256(token) → identity. Tokens are never stored in cleartext.
|
pub token: String,
|
||||||
hashed_tokens: Vec<([u8; 32], UserIdentity)>,
|
|
||||||
/// Original first token kept only for single-user startup printing.
|
|
||||||
/// Not used for authentication.
|
|
||||||
display_token: Option<String>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl MultiAuthState {
|
|
||||||
/// Create a single-user auth state (backwards compatible).
|
|
||||||
pub fn single(token: String, user_id: String) -> Self {
|
|
||||||
let hash = hash_token(&token);
|
|
||||||
Self {
|
|
||||||
hashed_tokens: vec![(
|
|
||||||
hash,
|
|
||||||
UserIdentity {
|
|
||||||
user_id,
|
|
||||||
workspace_read_scopes: Vec::new(),
|
|
||||||
},
|
|
||||||
)],
|
|
||||||
display_token: Some(token),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Create a multi-user auth state from a map of tokens to identities.
|
|
||||||
pub fn multi(tokens: HashMap<String, UserIdentity>) -> Self {
|
|
||||||
let hashed_tokens: Vec<([u8; 32], UserIdentity)> = tokens
|
|
||||||
.into_iter()
|
|
||||||
.map(|(tok, identity)| (hash_token(&tok), identity))
|
|
||||||
.collect();
|
|
||||||
Self {
|
|
||||||
hashed_tokens,
|
|
||||||
display_token: None,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Authenticate a token, returning the associated identity if valid.
|
|
||||||
///
|
|
||||||
/// Uses SHA-256 hashing + constant-time comparison (`subtle::ConstantTimeEq`)
|
|
||||||
/// to prevent timing side-channels. Both the candidate and stored tokens are
|
|
||||||
/// hashed to 32-byte digests, eliminating length-oracle leaks. Iterates all
|
|
||||||
/// entries regardless of match to avoid early-exit timing differences.
|
|
||||||
/// O(n) in the number of configured users — negligible for typical
|
|
||||||
/// deployments (< 10 users).
|
|
||||||
pub fn authenticate(&self, candidate: &str) -> Option<&UserIdentity> {
|
|
||||||
let candidate_hash = hash_token(candidate);
|
|
||||||
let mut matched: Option<&UserIdentity> = None;
|
|
||||||
for (stored_hash, identity) in &self.hashed_tokens {
|
|
||||||
if bool::from(candidate_hash.ct_eq(stored_hash)) {
|
|
||||||
matched = Some(identity);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
matched
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Get the first token for backwards-compatible printing at startup.
|
|
||||||
///
|
|
||||||
/// Only available in single-user mode; returns `None` in multi-user mode
|
|
||||||
/// to avoid exposing tokens.
|
|
||||||
pub fn first_token(&self) -> Option<&str> {
|
|
||||||
self.display_token.as_deref()
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Get the first user identity (for single-user fallback).
|
|
||||||
pub fn first_identity(&self) -> Option<&UserIdentity> {
|
|
||||||
self.hashed_tokens.first().map(|(_, id)| id)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Axum extractor that provides the authenticated user identity.
|
|
||||||
///
|
|
||||||
/// Only available on routes behind `auth_middleware`. Extracts the
|
|
||||||
/// `UserIdentity` that the middleware inserted into request extensions.
|
|
||||||
pub struct AuthenticatedUser(pub UserIdentity);
|
|
||||||
|
|
||||||
impl<S> FromRequestParts<S> for AuthenticatedUser
|
|
||||||
where
|
|
||||||
S: Send + Sync,
|
|
||||||
{
|
|
||||||
type Rejection = (StatusCode, &'static str);
|
|
||||||
|
|
||||||
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
|
|
||||||
parts
|
|
||||||
.extensions
|
|
||||||
.get::<UserIdentity>()
|
|
||||||
.cloned()
|
|
||||||
.map(AuthenticatedUser)
|
|
||||||
.ok_or((StatusCode::UNAUTHORIZED, "Not authenticated"))
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Whether query-string token auth is allowed for this request.
|
/// Whether query-string token auth is allowed for this request.
|
||||||
@@ -167,34 +51,29 @@ fn query_token(request: &Request) -> Option<String> {
|
|||||||
/// Auth middleware that validates bearer token from header or query param.
|
/// Auth middleware that validates bearer token from header or query param.
|
||||||
///
|
///
|
||||||
/// SSE connections can't set headers from `EventSource`, so we also accept
|
/// SSE connections can't set headers from `EventSource`, so we also accept
|
||||||
/// `?token=xxx` as a query parameter, but only on SSE/WS endpoints.
|
/// `?token=xxx` as a query parameter, but only on SSE endpoints.
|
||||||
///
|
|
||||||
/// On successful authentication, inserts the matching `UserIdentity` into
|
|
||||||
/// request extensions for downstream extraction via `AuthenticatedUser`.
|
|
||||||
pub async fn auth_middleware(
|
pub async fn auth_middleware(
|
||||||
State(auth): State<MultiAuthState>,
|
State(auth): State<AuthState>,
|
||||||
headers: HeaderMap,
|
headers: HeaderMap,
|
||||||
mut request: Request,
|
request: Request,
|
||||||
next: Next,
|
next: Next,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
// Try Authorization header first.
|
// Try Authorization header first (constant-time comparison).
|
||||||
// RFC 6750 Section 2.1: auth-scheme comparison is case-insensitive.
|
// RFC 6750 Section 2.1: auth-scheme comparison is case-insensitive.
|
||||||
if let Some(auth_header) = headers.get("authorization")
|
if let Some(auth_header) = headers.get("authorization")
|
||||||
&& let Ok(value) = auth_header.to_str()
|
&& let Ok(value) = auth_header.to_str()
|
||||||
&& value.len() > 7
|
&& value.len() > 7
|
||||||
&& value[..7].eq_ignore_ascii_case("Bearer ")
|
&& value[..7].eq_ignore_ascii_case("Bearer ")
|
||||||
&& let Some(identity) = auth.authenticate(&value[7..])
|
&& bool::from(value.as_bytes()[7..].ct_eq(auth.token.as_bytes()))
|
||||||
{
|
{
|
||||||
request.extensions_mut().insert(identity.clone());
|
|
||||||
return next.run(request).await;
|
return next.run(request).await;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Fall back to query parameter, but only for SSE/WS endpoints.
|
// Fall back to query parameter, but only for SSE endpoints (constant-time comparison).
|
||||||
if allows_query_token_auth(&request)
|
if allows_query_token_auth(&request)
|
||||||
&& let Some(token) = query_token(&request)
|
&& let Some(token) = query_token(&request)
|
||||||
&& let Some(identity) = auth.authenticate(&token)
|
&& bool::from(token.as_bytes().ct_eq(auth.token.as_bytes()))
|
||||||
{
|
{
|
||||||
request.extensions_mut().insert(identity.clone());
|
|
||||||
return next.run(request).await;
|
return next.run(request).await;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -204,61 +83,15 @@ pub async fn auth_middleware(
|
|||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
use crate::testing::credentials::TEST_AUTH_SECRET_TOKEN;
|
use crate::testing::credentials::{TEST_AUTH_SECRET_TOKEN, TEST_BEARER_TOKEN};
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_multi_auth_state_single() {
|
fn test_auth_state_clone() {
|
||||||
let state = MultiAuthState::single("tok-123".to_string(), "alice".to_string());
|
let state = AuthState {
|
||||||
let identity = state.authenticate("tok-123");
|
token: TEST_BEARER_TOKEN.to_string(),
|
||||||
assert!(identity.is_some());
|
};
|
||||||
assert_eq!(identity.unwrap().user_id, "alice");
|
let cloned = state.clone();
|
||||||
}
|
assert_eq!(cloned.token, TEST_BEARER_TOKEN);
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_multi_auth_state_reject_wrong_token() {
|
|
||||||
let state = MultiAuthState::single("tok-123".to_string(), "alice".to_string());
|
|
||||||
assert!(state.authenticate("wrong-token").is_none());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_multi_auth_state_multi_users() {
|
|
||||||
let mut tokens = HashMap::new();
|
|
||||||
tokens.insert(
|
|
||||||
"tok-alice".to_string(),
|
|
||||||
UserIdentity {
|
|
||||||
user_id: "alice".to_string(),
|
|
||||||
workspace_read_scopes: Vec::new(),
|
|
||||||
},
|
|
||||||
);
|
|
||||||
tokens.insert(
|
|
||||||
"tok-bob".to_string(),
|
|
||||||
UserIdentity {
|
|
||||||
user_id: "bob".to_string(),
|
|
||||||
workspace_read_scopes: Vec::new(),
|
|
||||||
},
|
|
||||||
);
|
|
||||||
let state = MultiAuthState::multi(tokens);
|
|
||||||
|
|
||||||
let alice = state.authenticate("tok-alice").unwrap();
|
|
||||||
assert_eq!(alice.user_id, "alice");
|
|
||||||
|
|
||||||
let bob = state.authenticate("tok-bob").unwrap();
|
|
||||||
assert_eq!(bob.user_id, "bob");
|
|
||||||
|
|
||||||
assert!(state.authenticate("tok-charlie").is_none());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_multi_auth_state_first_token() {
|
|
||||||
let state = MultiAuthState::single("my-token".to_string(), "user1".to_string());
|
|
||||||
assert_eq!(state.first_token(), Some("my-token"));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_multi_auth_state_first_identity() {
|
|
||||||
let state = MultiAuthState::single("my-token".to_string(), "user1".to_string());
|
|
||||||
let identity = state.first_identity().unwrap();
|
|
||||||
assert_eq!(identity.user_id, "user1");
|
|
||||||
}
|
}
|
||||||
|
|
||||||
use axum::Router;
|
use axum::Router;
|
||||||
@@ -274,7 +107,9 @@ mod tests {
|
|||||||
/// Router with streaming endpoints (query auth allowed) and regular
|
/// Router with streaming endpoints (query auth allowed) and regular
|
||||||
/// endpoints (query auth rejected).
|
/// endpoints (query auth rejected).
|
||||||
fn test_app(token: &str) -> Router {
|
fn test_app(token: &str) -> Router {
|
||||||
let state = MultiAuthState::single(token.to_string(), "test-user".to_string());
|
let state = AuthState {
|
||||||
|
token: token.to_string(),
|
||||||
|
};
|
||||||
Router::new()
|
Router::new()
|
||||||
.route("/api/chat/events", get(dummy_handler))
|
.route("/api/chat/events", get(dummy_handler))
|
||||||
.route("/api/logs/events", get(dummy_handler))
|
.route("/api/logs/events", get(dummy_handler))
|
||||||
@@ -471,200 +306,4 @@ mod tests {
|
|||||||
let resp = app.oneshot(req).await.unwrap();
|
let resp = app.oneshot(req).await.unwrap();
|
||||||
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
|
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
|
||||||
}
|
}
|
||||||
|
|
||||||
// --- Multi-tenant auth integration tests ---
|
|
||||||
|
|
||||||
/// Handler that extracts `AuthenticatedUser` and returns the resolved user_id.
|
|
||||||
async fn identity_handler(AuthenticatedUser(identity): AuthenticatedUser) -> String {
|
|
||||||
identity.user_id
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Handler that extracts `AuthenticatedUser` and returns workspace_read_scopes as JSON.
|
|
||||||
async fn scopes_handler(AuthenticatedUser(identity): AuthenticatedUser) -> String {
|
|
||||||
serde_json::to_string(&identity.workspace_read_scopes).unwrap()
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Build a multi-user router where each token maps to a distinct identity.
|
|
||||||
fn multi_user_app(tokens: HashMap<String, UserIdentity>) -> Router {
|
|
||||||
let state = MultiAuthState::multi(tokens);
|
|
||||||
Router::new()
|
|
||||||
.route("/api/chat/events", get(identity_handler))
|
|
||||||
.route("/api/chat/send", post(identity_handler))
|
|
||||||
.route("/api/scopes", get(scopes_handler))
|
|
||||||
.layer(middleware::from_fn_with_state(state, auth_middleware))
|
|
||||||
}
|
|
||||||
|
|
||||||
fn two_user_tokens() -> HashMap<String, UserIdentity> {
|
|
||||||
let mut tokens = HashMap::new();
|
|
||||||
tokens.insert(
|
|
||||||
"tok-alice".to_string(),
|
|
||||||
UserIdentity {
|
|
||||||
user_id: "alice".to_string(),
|
|
||||||
workspace_read_scopes: vec!["shared".to_string()],
|
|
||||||
},
|
|
||||||
);
|
|
||||||
tokens.insert(
|
|
||||||
"tok-bob".to_string(),
|
|
||||||
UserIdentity {
|
|
||||||
user_id: "bob".to_string(),
|
|
||||||
workspace_read_scopes: vec!["shared".to_string(), "alice".to_string()],
|
|
||||||
},
|
|
||||||
);
|
|
||||||
tokens
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_multi_user_alice_token_resolves_to_alice() {
|
|
||||||
let app = multi_user_app(two_user_tokens());
|
|
||||||
let req = Request::builder()
|
|
||||||
.uri("/api/chat/events")
|
|
||||||
.header("Authorization", "Bearer tok-alice")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(resp.status(), StatusCode::OK);
|
|
||||||
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
|
|
||||||
assert_eq!(body, "alice");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_multi_user_bob_token_resolves_to_bob() {
|
|
||||||
let app = multi_user_app(two_user_tokens());
|
|
||||||
let req = Request::builder()
|
|
||||||
.uri("/api/chat/events")
|
|
||||||
.header("Authorization", "Bearer tok-bob")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(resp.status(), StatusCode::OK);
|
|
||||||
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
|
|
||||||
assert_eq!(body, "bob");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_multi_user_sequential_tokens_resolve_independently() {
|
|
||||||
// Send both alice and bob tokens sequentially and verify each gets
|
|
||||||
// the correct identity — guards against token map corruption.
|
|
||||||
let tokens = two_user_tokens();
|
|
||||||
|
|
||||||
let app1 = multi_user_app(tokens.clone());
|
|
||||||
let req = Request::builder()
|
|
||||||
.uri("/api/chat/events")
|
|
||||||
.header("Authorization", "Bearer tok-alice")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app1.oneshot(req).await.unwrap();
|
|
||||||
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
|
|
||||||
assert_eq!(body, "alice");
|
|
||||||
|
|
||||||
let app2 = multi_user_app(tokens);
|
|
||||||
let req = Request::builder()
|
|
||||||
.uri("/api/chat/events")
|
|
||||||
.header("Authorization", "Bearer tok-bob")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app2.oneshot(req).await.unwrap();
|
|
||||||
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
|
|
||||||
assert_eq!(body, "bob");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_multi_user_unknown_token_rejected() {
|
|
||||||
let app = multi_user_app(two_user_tokens());
|
|
||||||
let req = Request::builder()
|
|
||||||
.uri("/api/chat/events")
|
|
||||||
.header("Authorization", "Bearer tok-charlie")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_multi_user_workspace_read_scopes_propagated() {
|
|
||||||
let app = multi_user_app(two_user_tokens());
|
|
||||||
|
|
||||||
// Alice has ["shared"]
|
|
||||||
let req = Request::builder()
|
|
||||||
.uri("/api/scopes")
|
|
||||||
.header("Authorization", "Bearer tok-alice")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.oneshot(req).await.unwrap();
|
|
||||||
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
|
|
||||||
let scopes: Vec<String> = serde_json::from_slice(&body).unwrap();
|
|
||||||
assert_eq!(scopes, vec!["shared"]);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_multi_user_bob_has_two_scopes() {
|
|
||||||
let app = multi_user_app(two_user_tokens());
|
|
||||||
|
|
||||||
// Bob has ["shared", "alice"]
|
|
||||||
let req = Request::builder()
|
|
||||||
.uri("/api/scopes")
|
|
||||||
.header("Authorization", "Bearer tok-bob")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.oneshot(req).await.unwrap();
|
|
||||||
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
|
|
||||||
let scopes: Vec<String> = serde_json::from_slice(&body).unwrap();
|
|
||||||
assert_eq!(scopes, vec!["shared", "alice"]);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_multi_user_query_param_resolves_correct_identity() {
|
|
||||||
let app = multi_user_app(two_user_tokens());
|
|
||||||
let req = Request::builder()
|
|
||||||
.uri("/api/chat/events?token=tok-bob")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(resp.status(), StatusCode::OK);
|
|
||||||
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
|
|
||||||
assert_eq!(body, "bob");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_multi_user_post_with_bearer_resolves_identity() {
|
|
||||||
let app = multi_user_app(two_user_tokens());
|
|
||||||
let req = Request::builder()
|
|
||||||
.method(Method::POST)
|
|
||||||
.uri("/api/chat/send")
|
|
||||||
.header("Authorization", "Bearer tok-alice")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(resp.status(), StatusCode::OK);
|
|
||||||
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
|
|
||||||
assert_eq!(body, "alice");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_multi_user_empty_scopes_for_single_user() {
|
|
||||||
// Single-user mode creates identity with empty workspace_read_scopes.
|
|
||||||
let state = MultiAuthState::single("tok-only".to_string(), "solo".to_string());
|
|
||||||
let app = Router::new()
|
|
||||||
.route("/api/scopes", get(scopes_handler))
|
|
||||||
.layer(middleware::from_fn_with_state(state, auth_middleware));
|
|
||||||
let req = Request::builder()
|
|
||||||
.uri("/api/scopes")
|
|
||||||
.header("Authorization", "Bearer tok-only")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.oneshot(req).await.unwrap();
|
|
||||||
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
|
|
||||||
let scopes: Vec<String> = serde_json::from_slice(&body).unwrap();
|
|
||||||
assert!(scopes.is_empty());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_prefix_and_extension_tokens_rejected() {
|
|
||||||
// Verifies that prefix/suffix variants of valid tokens are rejected.
|
|
||||||
// Note: the constant-time property is enforced structurally by use of
|
|
||||||
// subtle::ConstantTimeEq and cannot be verified via outcome testing.
|
|
||||||
let state = MultiAuthState::single("long-secret-token".to_string(), "user".to_string());
|
|
||||||
assert!(state.authenticate("long-secret").is_none());
|
|
||||||
assert!(state.authenticate("long-secret-token-extra").is_none());
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -12,24 +12,22 @@ use serde::Deserialize;
|
|||||||
use uuid::Uuid;
|
use uuid::Uuid;
|
||||||
|
|
||||||
use crate::channels::IncomingMessage;
|
use crate::channels::IncomingMessage;
|
||||||
use crate::channels::web::auth::AuthenticatedUser;
|
|
||||||
use crate::channels::web::server::GatewayState;
|
use crate::channels::web::server::GatewayState;
|
||||||
use crate::channels::web::types::*;
|
use crate::channels::web::types::*;
|
||||||
use crate::channels::web::util::{build_turns_from_db_messages, truncate_preview};
|
use crate::channels::web::util::{build_turns_from_db_messages, truncate_preview};
|
||||||
|
|
||||||
pub async fn chat_send_handler(
|
pub async fn chat_send_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(identity): AuthenticatedUser,
|
|
||||||
Json(req): Json<SendMessageRequest>,
|
Json(req): Json<SendMessageRequest>,
|
||||||
) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> {
|
) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> {
|
||||||
if !state.chat_rate_limiter.check(&identity.user_id) {
|
if !state.chat_rate_limiter.check() {
|
||||||
return Err((
|
return Err((
|
||||||
StatusCode::TOO_MANY_REQUESTS,
|
StatusCode::TOO_MANY_REQUESTS,
|
||||||
"Rate limit exceeded. Try again shortly.".to_string(),
|
"Rate limit exceeded. Try again shortly.".to_string(),
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
let mut msg = IncomingMessage::new("gateway", &identity.user_id, &req.content);
|
let mut msg = IncomingMessage::new("gateway", &state.user_id, &req.content);
|
||||||
|
|
||||||
if let Some(ref thread_id) = req.thread_id {
|
if let Some(ref thread_id) = req.thread_id {
|
||||||
msg = msg.with_thread(thread_id);
|
msg = msg.with_thread(thread_id);
|
||||||
@@ -76,7 +74,6 @@ pub async fn chat_send_handler(
|
|||||||
|
|
||||||
pub async fn chat_approval_handler(
|
pub async fn chat_approval_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(identity): AuthenticatedUser,
|
|
||||||
Json(req): Json<ApprovalRequest>,
|
Json(req): Json<ApprovalRequest>,
|
||||||
) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> {
|
) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> {
|
||||||
let (approved, always) = match req.action.as_str() {
|
let (approved, always) = match req.action.as_str() {
|
||||||
@@ -112,7 +109,7 @@ pub async fn chat_approval_handler(
|
|||||||
)
|
)
|
||||||
})?;
|
})?;
|
||||||
|
|
||||||
let mut msg = IncomingMessage::new("gateway", &identity.user_id, content);
|
let mut msg = IncomingMessage::new("gateway", &state.user_id, content);
|
||||||
|
|
||||||
if let Some(ref thread_id) = req.thread_id {
|
if let Some(ref thread_id) = req.thread_id {
|
||||||
msg = msg.with_thread(thread_id);
|
msg = msg.with_thread(thread_id);
|
||||||
@@ -153,7 +150,6 @@ pub async fn chat_approval_handler(
|
|||||||
/// The token never touches the LLM, chat history, or SSE stream.
|
/// The token never touches the LLM, chat history, or SSE stream.
|
||||||
pub async fn chat_auth_token_handler(
|
pub async fn chat_auth_token_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Json(req): Json<AuthTokenRequest>,
|
Json(req): Json<AuthTokenRequest>,
|
||||||
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
|
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
|
||||||
let ext_mgr = state.extension_manager.as_ref().ok_or((
|
let ext_mgr = state.extension_manager.as_ref().ok_or((
|
||||||
@@ -162,7 +158,7 @@ pub async fn chat_auth_token_handler(
|
|||||||
))?;
|
))?;
|
||||||
|
|
||||||
match ext_mgr
|
match ext_mgr
|
||||||
.configure_token(&req.extension_name, &req.token, &user.user_id)
|
.configure_token(&req.extension_name, &req.token)
|
||||||
.await
|
.await
|
||||||
{
|
{
|
||||||
Ok(result) => {
|
Ok(result) => {
|
||||||
@@ -173,26 +169,20 @@ pub async fn chat_auth_token_handler(
|
|||||||
resp.instructions = result.verification.as_ref().map(|v| v.instructions.clone());
|
resp.instructions = result.verification.as_ref().map(|v| v.instructions.clone());
|
||||||
|
|
||||||
if result.verification.is_some() {
|
if result.verification.is_some() {
|
||||||
state.sse.broadcast_for_user(
|
state.sse.broadcast(SseEvent::AuthRequired {
|
||||||
&user.user_id,
|
extension_name: req.extension_name.clone(),
|
||||||
SseEvent::AuthRequired {
|
instructions: Some(result.message),
|
||||||
extension_name: req.extension_name.clone(),
|
auth_url: None,
|
||||||
instructions: Some(result.message),
|
setup_url: None,
|
||||||
auth_url: None,
|
});
|
||||||
setup_url: None,
|
|
||||||
},
|
|
||||||
);
|
|
||||||
} else {
|
} else {
|
||||||
clear_auth_mode(&state, &user.user_id).await;
|
clear_auth_mode(&state).await;
|
||||||
|
|
||||||
state.sse.broadcast_for_user(
|
state.sse.broadcast(SseEvent::AuthCompleted {
|
||||||
&user.user_id,
|
extension_name: req.extension_name.clone(),
|
||||||
SseEvent::AuthCompleted {
|
success: true,
|
||||||
extension_name: req.extension_name.clone(),
|
message: result.message,
|
||||||
success: true,
|
});
|
||||||
message: result.message,
|
|
||||||
},
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(Json(resp))
|
Ok(Json(resp))
|
||||||
@@ -200,15 +190,12 @@ pub async fn chat_auth_token_handler(
|
|||||||
Err(e) => {
|
Err(e) => {
|
||||||
let msg = e.to_string();
|
let msg = e.to_string();
|
||||||
if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) {
|
if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) {
|
||||||
state.sse.broadcast_for_user(
|
state.sse.broadcast(SseEvent::AuthRequired {
|
||||||
&user.user_id,
|
extension_name: req.extension_name.clone(),
|
||||||
SseEvent::AuthRequired {
|
instructions: Some(msg.clone()),
|
||||||
extension_name: req.extension_name.clone(),
|
auth_url: None,
|
||||||
instructions: Some(msg.clone()),
|
setup_url: None,
|
||||||
auth_url: None,
|
});
|
||||||
setup_url: None,
|
|
||||||
},
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
Ok(Json(ActionResponse::fail(msg)))
|
Ok(Json(ActionResponse::fail(msg)))
|
||||||
}
|
}
|
||||||
@@ -218,17 +205,16 @@ pub async fn chat_auth_token_handler(
|
|||||||
/// Cancel an in-progress auth flow.
|
/// Cancel an in-progress auth flow.
|
||||||
pub async fn chat_auth_cancel_handler(
|
pub async fn chat_auth_cancel_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(identity): AuthenticatedUser,
|
|
||||||
Json(_req): Json<AuthCancelRequest>,
|
Json(_req): Json<AuthCancelRequest>,
|
||||||
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
|
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
|
||||||
clear_auth_mode(&state, &identity.user_id).await;
|
clear_auth_mode(&state).await;
|
||||||
Ok(Json(ActionResponse::ok("Auth cancelled")))
|
Ok(Json(ActionResponse::ok("Auth cancelled")))
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Clear pending auth mode on the active thread.
|
/// Clear pending auth mode on the active thread.
|
||||||
pub async fn clear_auth_mode(state: &GatewayState, user_id: &str) {
|
pub async fn clear_auth_mode(state: &GatewayState) {
|
||||||
if let Some(ref sm) = state.session_manager {
|
if let Some(ref sm) = state.session_manager {
|
||||||
let session = sm.get_or_create_session(user_id).await;
|
let session = sm.get_or_create_session(&state.user_id).await;
|
||||||
let mut sess = session.lock().await;
|
let mut sess = session.lock().await;
|
||||||
if let Some(thread_id) = sess.active_thread
|
if let Some(thread_id) = sess.active_thread
|
||||||
&& let Some(thread) = sess.threads.get_mut(&thread_id)
|
&& let Some(thread) = sess.threads.get_mut(&thread_id)
|
||||||
@@ -240,9 +226,8 @@ pub async fn clear_auth_mode(state: &GatewayState, user_id: &str) {
|
|||||||
|
|
||||||
pub async fn chat_events_handler(
|
pub async fn chat_events_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
) -> Result<impl IntoResponse, (StatusCode, String)> {
|
) -> Result<impl IntoResponse, (StatusCode, String)> {
|
||||||
state.sse.subscribe(Some(user.user_id)).ok_or((
|
state.sse.subscribe().ok_or((
|
||||||
StatusCode::SERVICE_UNAVAILABLE,
|
StatusCode::SERVICE_UNAVAILABLE,
|
||||||
"Too many connections".to_string(),
|
"Too many connections".to_string(),
|
||||||
))
|
))
|
||||||
@@ -252,7 +237,6 @@ pub async fn chat_ws_handler(
|
|||||||
headers: axum::http::HeaderMap,
|
headers: axum::http::HeaderMap,
|
||||||
ws: WebSocketUpgrade,
|
ws: WebSocketUpgrade,
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(identity): AuthenticatedUser,
|
|
||||||
) -> Result<impl IntoResponse, (StatusCode, String)> {
|
) -> Result<impl IntoResponse, (StatusCode, String)> {
|
||||||
// Validate Origin header to prevent cross-site WebSocket hijacking.
|
// Validate Origin header to prevent cross-site WebSocket hijacking.
|
||||||
let origin = headers
|
let origin = headers
|
||||||
@@ -278,9 +262,7 @@ pub async fn chat_ws_handler(
|
|||||||
"WebSocket origin not allowed".to_string(),
|
"WebSocket origin not allowed".to_string(),
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
Ok(ws.on_upgrade(move |socket| {
|
Ok(ws.on_upgrade(move |socket| crate::channels::web::ws::handle_ws_connection(socket, state)))
|
||||||
crate::channels::web::ws::handle_ws_connection(socket, state, identity)
|
|
||||||
}))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Deserialize)]
|
#[derive(Deserialize)]
|
||||||
@@ -292,7 +274,6 @@ pub struct HistoryQuery {
|
|||||||
|
|
||||||
pub async fn chat_history_handler(
|
pub async fn chat_history_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(identity): AuthenticatedUser,
|
|
||||||
Query(query): Query<HistoryQuery>,
|
Query(query): Query<HistoryQuery>,
|
||||||
) -> Result<Json<HistoryResponse>, (StatusCode, String)> {
|
) -> Result<Json<HistoryResponse>, (StatusCode, String)> {
|
||||||
let session_manager = state.session_manager.as_ref().ok_or((
|
let session_manager = state.session_manager.as_ref().ok_or((
|
||||||
@@ -300,9 +281,7 @@ pub async fn chat_history_handler(
|
|||||||
"Session manager not available".to_string(),
|
"Session manager not available".to_string(),
|
||||||
))?;
|
))?;
|
||||||
|
|
||||||
let session = session_manager
|
let session = session_manager.get_or_create_session(&state.user_id).await;
|
||||||
.get_or_create_session(&identity.user_id)
|
|
||||||
.await;
|
|
||||||
|
|
||||||
let limit = query.limit.unwrap_or(50);
|
let limit = query.limit.unwrap_or(50);
|
||||||
let before_cursor = query
|
let before_cursor = query
|
||||||
@@ -335,7 +314,7 @@ pub async fn chat_history_handler(
|
|||||||
&& let Some(ref store) = state.store
|
&& let Some(ref store) = state.store
|
||||||
{
|
{
|
||||||
let owned = store
|
let owned = store
|
||||||
.conversation_belongs_to_user(thread_id, &identity.user_id)
|
.conversation_belongs_to_user(thread_id, &state.user_id)
|
||||||
.await
|
.await
|
||||||
.unwrap_or(false);
|
.unwrap_or(false);
|
||||||
if !owned {
|
if !owned {
|
||||||
@@ -455,27 +434,24 @@ pub async fn chat_history_handler(
|
|||||||
|
|
||||||
pub async fn chat_threads_handler(
|
pub async fn chat_threads_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(identity): AuthenticatedUser,
|
|
||||||
) -> Result<Json<ThreadListResponse>, (StatusCode, String)> {
|
) -> Result<Json<ThreadListResponse>, (StatusCode, String)> {
|
||||||
let session_manager = state.session_manager.as_ref().ok_or((
|
let session_manager = state.session_manager.as_ref().ok_or((
|
||||||
StatusCode::SERVICE_UNAVAILABLE,
|
StatusCode::SERVICE_UNAVAILABLE,
|
||||||
"Session manager not available".to_string(),
|
"Session manager not available".to_string(),
|
||||||
))?;
|
))?;
|
||||||
|
|
||||||
let session = session_manager
|
let session = session_manager.get_or_create_session(&state.user_id).await;
|
||||||
.get_or_create_session(&identity.user_id)
|
|
||||||
.await;
|
|
||||||
|
|
||||||
// Try DB first for persistent thread list
|
// Try DB first for persistent thread list
|
||||||
if let Some(ref store) = state.store {
|
if let Some(ref store) = state.store {
|
||||||
// Auto-create assistant thread if it doesn't exist
|
// Auto-create assistant thread if it doesn't exist
|
||||||
let assistant_id = store
|
let assistant_id = store
|
||||||
.get_or_create_assistant_conversation(&identity.user_id, "gateway")
|
.get_or_create_assistant_conversation(&state.user_id, "gateway")
|
||||||
.await
|
.await
|
||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||||
|
|
||||||
if let Ok(summaries) = store
|
if let Ok(summaries) = store
|
||||||
.list_conversations_all_channels(&identity.user_id, 50)
|
.list_conversations_all_channels(&state.user_id, 50)
|
||||||
.await
|
.await
|
||||||
{
|
{
|
||||||
let mut assistant_thread = None;
|
let mut assistant_thread = None;
|
||||||
@@ -558,16 +534,13 @@ pub async fn chat_threads_handler(
|
|||||||
|
|
||||||
pub async fn chat_new_thread_handler(
|
pub async fn chat_new_thread_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(identity): AuthenticatedUser,
|
|
||||||
) -> Result<Json<ThreadInfo>, (StatusCode, String)> {
|
) -> Result<Json<ThreadInfo>, (StatusCode, String)> {
|
||||||
let session_manager = state.session_manager.as_ref().ok_or((
|
let session_manager = state.session_manager.as_ref().ok_or((
|
||||||
StatusCode::SERVICE_UNAVAILABLE,
|
StatusCode::SERVICE_UNAVAILABLE,
|
||||||
"Session manager not available".to_string(),
|
"Session manager not available".to_string(),
|
||||||
))?;
|
))?;
|
||||||
|
|
||||||
let session = session_manager
|
let session = session_manager.get_or_create_session(&state.user_id).await;
|
||||||
.get_or_create_session(&identity.user_id)
|
|
||||||
.await;
|
|
||||||
let (thread_id, info) = {
|
let (thread_id, info) = {
|
||||||
let mut sess = session.lock().await;
|
let mut sess = session.lock().await;
|
||||||
let thread = sess.create_thread();
|
let thread = sess.create_thread();
|
||||||
@@ -589,12 +562,12 @@ pub async fn chat_new_thread_handler(
|
|||||||
// so that the subsequent loadThreads() call from the frontend sees it.
|
// so that the subsequent loadThreads() call from the frontend sees it.
|
||||||
if let Some(ref store) = state.store {
|
if let Some(ref store) = state.store {
|
||||||
match store
|
match store
|
||||||
.ensure_conversation(thread_id, "gateway", &identity.user_id, None)
|
.ensure_conversation(thread_id, "gateway", &state.user_id, None)
|
||||||
.await
|
.await
|
||||||
{
|
{
|
||||||
Ok(true) => {}
|
Ok(true) => {}
|
||||||
Ok(false) => tracing::warn!(
|
Ok(false) => tracing::warn!(
|
||||||
user = %identity.user_id,
|
user = %state.user_id,
|
||||||
thread_id = %thread_id,
|
thread_id = %thread_id,
|
||||||
"Skipped persisting new thread due to ownership/channel conflict"
|
"Skipped persisting new thread due to ownership/channel conflict"
|
||||||
),
|
),
|
||||||
|
|||||||
@@ -8,13 +8,11 @@ use axum::{
|
|||||||
http::StatusCode,
|
http::StatusCode,
|
||||||
};
|
};
|
||||||
|
|
||||||
use crate::channels::web::auth::AuthenticatedUser;
|
|
||||||
use crate::channels::web::server::GatewayState;
|
use crate::channels::web::server::GatewayState;
|
||||||
use crate::channels::web::types::*;
|
use crate::channels::web::types::*;
|
||||||
|
|
||||||
pub async fn extensions_list_handler(
|
pub async fn extensions_list_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
) -> Result<Json<ExtensionListResponse>, (StatusCode, String)> {
|
) -> Result<Json<ExtensionListResponse>, (StatusCode, String)> {
|
||||||
let ext_mgr = state.extension_manager.as_ref().ok_or((
|
let ext_mgr = state.extension_manager.as_ref().ok_or((
|
||||||
StatusCode::NOT_IMPLEMENTED,
|
StatusCode::NOT_IMPLEMENTED,
|
||||||
@@ -22,7 +20,7 @@ pub async fn extensions_list_handler(
|
|||||||
))?;
|
))?;
|
||||||
|
|
||||||
let installed = ext_mgr
|
let installed = ext_mgr
|
||||||
.list(None, false, &user.user_id)
|
.list(None, false)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||||
|
|
||||||
@@ -82,7 +80,6 @@ pub async fn extensions_list_handler(
|
|||||||
|
|
||||||
pub async fn extensions_tools_handler(
|
pub async fn extensions_tools_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(_user): AuthenticatedUser,
|
|
||||||
) -> Result<Json<ToolListResponse>, (StatusCode, String)> {
|
) -> Result<Json<ToolListResponse>, (StatusCode, String)> {
|
||||||
let registry = state.tool_registry.as_ref().ok_or((
|
let registry = state.tool_registry.as_ref().ok_or((
|
||||||
StatusCode::SERVICE_UNAVAILABLE,
|
StatusCode::SERVICE_UNAVAILABLE,
|
||||||
@@ -103,7 +100,6 @@ pub async fn extensions_tools_handler(
|
|||||||
|
|
||||||
pub async fn extensions_install_handler(
|
pub async fn extensions_install_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Json(req): Json<InstallExtensionRequest>,
|
Json(req): Json<InstallExtensionRequest>,
|
||||||
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
|
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
|
||||||
let ext_mgr = state.extension_manager.as_ref().ok_or((
|
let ext_mgr = state.extension_manager.as_ref().ok_or((
|
||||||
@@ -120,7 +116,7 @@ pub async fn extensions_install_handler(
|
|||||||
});
|
});
|
||||||
|
|
||||||
match ext_mgr
|
match ext_mgr
|
||||||
.install(&req.name, req.url.as_deref(), kind_hint, &user.user_id)
|
.install(&req.name, req.url.as_deref(), kind_hint)
|
||||||
.await
|
.await
|
||||||
{
|
{
|
||||||
Ok(result) => Ok(Json(ActionResponse::ok(result.message))),
|
Ok(result) => Ok(Json(ActionResponse::ok(result.message))),
|
||||||
@@ -130,7 +126,6 @@ pub async fn extensions_install_handler(
|
|||||||
|
|
||||||
pub async fn extensions_remove_handler(
|
pub async fn extensions_remove_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Path(name): Path<String>,
|
Path(name): Path<String>,
|
||||||
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
|
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
|
||||||
let ext_mgr = state.extension_manager.as_ref().ok_or((
|
let ext_mgr = state.extension_manager.as_ref().ok_or((
|
||||||
@@ -138,7 +133,7 @@ pub async fn extensions_remove_handler(
|
|||||||
"Extension manager not available (secrets store required)".to_string(),
|
"Extension manager not available (secrets store required)".to_string(),
|
||||||
))?;
|
))?;
|
||||||
|
|
||||||
match ext_mgr.remove(&name, &user.user_id).await {
|
match ext_mgr.remove(&name).await {
|
||||||
Ok(message) => Ok(Json(ActionResponse::ok(message))),
|
Ok(message) => Ok(Json(ActionResponse::ok(message))),
|
||||||
Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))),
|
Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))),
|
||||||
}
|
}
|
||||||
|
|||||||
+277
-388
@@ -11,13 +11,11 @@ use axum::{
|
|||||||
use serde::Deserialize;
|
use serde::Deserialize;
|
||||||
use uuid::Uuid;
|
use uuid::Uuid;
|
||||||
|
|
||||||
use crate::channels::web::auth::AuthenticatedUser;
|
|
||||||
use crate::channels::web::server::GatewayState;
|
use crate::channels::web::server::GatewayState;
|
||||||
use crate::channels::web::types::*;
|
use crate::channels::web::types::*;
|
||||||
|
|
||||||
pub async fn jobs_list_handler(
|
pub async fn jobs_list_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
) -> Result<Json<JobListResponse>, (StatusCode, String)> {
|
) -> Result<Json<JobListResponse>, (StatusCode, String)> {
|
||||||
let store = state.store.as_ref().ok_or((
|
let store = state.store.as_ref().ok_or((
|
||||||
StatusCode::SERVICE_UNAVAILABLE,
|
StatusCode::SERVICE_UNAVAILABLE,
|
||||||
@@ -27,8 +25,8 @@ pub async fn jobs_list_handler(
|
|||||||
let mut jobs: Vec<JobInfo> = Vec::new();
|
let mut jobs: Vec<JobInfo> = Vec::new();
|
||||||
let mut seen_ids: HashSet<Uuid> = HashSet::new();
|
let mut seen_ids: HashSet<Uuid> = HashSet::new();
|
||||||
|
|
||||||
// Fetch sandbox jobs scoped to this user.
|
// Fetch sandbox jobs from database.
|
||||||
match store.list_sandbox_jobs_for_user(&user.user_id).await {
|
match store.list_sandbox_jobs().await {
|
||||||
Ok(sandbox_jobs) => {
|
Ok(sandbox_jobs) => {
|
||||||
for j in &sandbox_jobs {
|
for j in &sandbox_jobs {
|
||||||
let ui_state = match j.status.as_str() {
|
let ui_state = match j.status.as_str() {
|
||||||
@@ -52,8 +50,8 @@ pub async fn jobs_list_handler(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Fetch agent (non-sandbox) jobs scoped to this user, deduplicating by ID.
|
// Fetch agent (non-sandbox) jobs from database, deduplicating by ID.
|
||||||
match store.list_agent_jobs_for_user(&user.user_id).await {
|
match store.list_agent_jobs().await {
|
||||||
Ok(agent_jobs) => {
|
Ok(agent_jobs) => {
|
||||||
for j in &agent_jobs {
|
for j in &agent_jobs {
|
||||||
if seen_ids.contains(&j.id) {
|
if seen_ids.contains(&j.id) {
|
||||||
@@ -82,7 +80,6 @@ pub async fn jobs_list_handler(
|
|||||||
|
|
||||||
pub async fn jobs_summary_handler(
|
pub async fn jobs_summary_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
) -> Result<Json<JobSummaryResponse>, (StatusCode, String)> {
|
) -> Result<Json<JobSummaryResponse>, (StatusCode, String)> {
|
||||||
let store = state.store.as_ref().ok_or((
|
let store = state.store.as_ref().ok_or((
|
||||||
StatusCode::SERVICE_UNAVAILABLE,
|
StatusCode::SERVICE_UNAVAILABLE,
|
||||||
@@ -96,8 +93,8 @@ pub async fn jobs_summary_handler(
|
|||||||
let mut failed = 0;
|
let mut failed = 0;
|
||||||
let mut stuck = 0;
|
let mut stuck = 0;
|
||||||
|
|
||||||
// Sandbox job counts scoped to this user.
|
// Sandbox job counts.
|
||||||
match store.sandbox_job_summary_for_user(&user.user_id).await {
|
match store.sandbox_job_summary().await {
|
||||||
Ok(s) => {
|
Ok(s) => {
|
||||||
total += s.total;
|
total += s.total;
|
||||||
pending += s.creating;
|
pending += s.creating;
|
||||||
@@ -110,8 +107,8 @@ pub async fn jobs_summary_handler(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Agent job counts scoped to this user.
|
// Agent job counts.
|
||||||
match store.agent_job_summary_for_user(&user.user_id).await {
|
match store.agent_job_summary().await {
|
||||||
Ok(s) => {
|
Ok(s) => {
|
||||||
total += s.total;
|
total += s.total;
|
||||||
pending += s.pending;
|
pending += s.pending;
|
||||||
@@ -137,7 +134,6 @@ pub async fn jobs_summary_handler(
|
|||||||
|
|
||||||
pub async fn jobs_detail_handler(
|
pub async fn jobs_detail_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Path(id): Path<String>,
|
Path(id): Path<String>,
|
||||||
) -> Result<Json<JobDetailResponse>, (StatusCode, String)> {
|
) -> Result<Json<JobDetailResponse>, (StatusCode, String)> {
|
||||||
let store = state.store.as_ref().ok_or((
|
let store = state.store.as_ref().ok_or((
|
||||||
@@ -149,213 +145,169 @@ pub async fn jobs_detail_handler(
|
|||||||
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
|
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
|
||||||
|
|
||||||
// Try sandbox job from DB first.
|
// Try sandbox job from DB first.
|
||||||
match store.get_sandbox_job(job_id).await {
|
if let Ok(Some(job)) = store.get_sandbox_job(job_id).await {
|
||||||
Ok(Some(job)) => {
|
let browse_id = std::path::Path::new(&job.project_dir)
|
||||||
if job.user_id != user.user_id {
|
.file_name()
|
||||||
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
|
.map(|n| n.to_string_lossy().to_string())
|
||||||
}
|
.unwrap_or_else(|| job.id.to_string());
|
||||||
let browse_id = std::path::Path::new(&job.project_dir)
|
|
||||||
.file_name()
|
|
||||||
.map(|n| n.to_string_lossy().to_string())
|
|
||||||
.unwrap_or_else(|| job.id.to_string());
|
|
||||||
|
|
||||||
let ui_state = match job.status.as_str() {
|
let ui_state = match job.status.as_str() {
|
||||||
"creating" => "pending",
|
"creating" => "pending",
|
||||||
"running" => "in_progress",
|
"running" => "in_progress",
|
||||||
s => s,
|
s => s,
|
||||||
};
|
};
|
||||||
|
|
||||||
let elapsed_secs = job.started_at.map(|start| {
|
let elapsed_secs = job.started_at.map(|start| {
|
||||||
let end = job.completed_at.unwrap_or_else(chrono::Utc::now);
|
let end = job.completed_at.unwrap_or_else(chrono::Utc::now);
|
||||||
(end - start).num_seconds().max(0) as u64
|
(end - start).num_seconds().max(0) as u64
|
||||||
|
});
|
||||||
|
|
||||||
|
// Synthesize transitions from timestamps.
|
||||||
|
let mut transitions = Vec::new();
|
||||||
|
if let Some(started) = job.started_at {
|
||||||
|
transitions.push(TransitionInfo {
|
||||||
|
from: "creating".to_string(),
|
||||||
|
to: "running".to_string(),
|
||||||
|
timestamp: started.to_rfc3339(),
|
||||||
|
reason: None,
|
||||||
});
|
});
|
||||||
|
|
||||||
// Synthesize transitions from timestamps.
|
|
||||||
let mut transitions = Vec::new();
|
|
||||||
if let Some(started) = job.started_at {
|
|
||||||
transitions.push(TransitionInfo {
|
|
||||||
from: "creating".to_string(),
|
|
||||||
to: "running".to_string(),
|
|
||||||
timestamp: started.to_rfc3339(),
|
|
||||||
reason: None,
|
|
||||||
});
|
|
||||||
}
|
|
||||||
if let Some(completed) = job.completed_at {
|
|
||||||
transitions.push(TransitionInfo {
|
|
||||||
from: "running".to_string(),
|
|
||||||
to: job.status.clone(),
|
|
||||||
timestamp: completed.to_rfc3339(),
|
|
||||||
reason: job.failure_reason.clone(),
|
|
||||||
});
|
|
||||||
}
|
|
||||||
|
|
||||||
let mode = store.get_sandbox_job_mode(job.id).await.ok().flatten();
|
|
||||||
let is_claude_code = mode.as_deref() == Some("claude_code");
|
|
||||||
|
|
||||||
return Ok(Json(JobDetailResponse {
|
|
||||||
id: job.id,
|
|
||||||
title: job.task.clone(),
|
|
||||||
description: String::new(),
|
|
||||||
state: ui_state.to_string(),
|
|
||||||
user_id: job.user_id.clone(),
|
|
||||||
created_at: job.created_at.to_rfc3339(),
|
|
||||||
started_at: job.started_at.map(|dt| dt.to_rfc3339()),
|
|
||||||
completed_at: job.completed_at.map(|dt| dt.to_rfc3339()),
|
|
||||||
elapsed_secs,
|
|
||||||
project_dir: Some(job.project_dir.clone()),
|
|
||||||
browse_url: Some(format!("/projects/{}/", browse_id)),
|
|
||||||
job_mode: mode.filter(|m| m != "worker"),
|
|
||||||
transitions,
|
|
||||||
can_restart: state.job_manager.is_some(),
|
|
||||||
can_prompt: is_claude_code && state.prompt_queue.is_some(),
|
|
||||||
job_kind: Some("sandbox".to_string()),
|
|
||||||
}));
|
|
||||||
}
|
}
|
||||||
Ok(None) => {}
|
if let Some(completed) = job.completed_at {
|
||||||
Err(e) => {
|
transitions.push(TransitionInfo {
|
||||||
return Err((
|
from: "running".to_string(),
|
||||||
StatusCode::INTERNAL_SERVER_ERROR,
|
to: job.status.clone(),
|
||||||
format!("Database error: {}", e),
|
timestamp: completed.to_rfc3339(),
|
||||||
));
|
reason: job.failure_reason.clone(),
|
||||||
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
|
let mode = store.get_sandbox_job_mode(job.id).await.ok().flatten();
|
||||||
|
let is_claude_code = mode.as_deref() == Some("claude_code");
|
||||||
|
|
||||||
|
return Ok(Json(JobDetailResponse {
|
||||||
|
id: job.id,
|
||||||
|
title: job.task.clone(),
|
||||||
|
description: String::new(),
|
||||||
|
state: ui_state.to_string(),
|
||||||
|
user_id: job.user_id.clone(),
|
||||||
|
created_at: job.created_at.to_rfc3339(),
|
||||||
|
started_at: job.started_at.map(|dt| dt.to_rfc3339()),
|
||||||
|
completed_at: job.completed_at.map(|dt| dt.to_rfc3339()),
|
||||||
|
elapsed_secs,
|
||||||
|
project_dir: Some(job.project_dir.clone()),
|
||||||
|
browse_url: Some(format!("/projects/{}/", browse_id)),
|
||||||
|
job_mode: mode.filter(|m| m != "worker"),
|
||||||
|
transitions,
|
||||||
|
can_restart: state.job_manager.is_some(),
|
||||||
|
can_prompt: is_claude_code && state.prompt_queue.is_some(),
|
||||||
|
job_kind: Some("sandbox".to_string()),
|
||||||
|
}));
|
||||||
}
|
}
|
||||||
|
|
||||||
// Fall back to agent job from DB.
|
// Fall back to agent job from DB.
|
||||||
match store.get_job(job_id).await {
|
if let Ok(Some(ctx)) = store.get_job(job_id).await {
|
||||||
Ok(Some(ctx)) => {
|
let elapsed_secs = ctx.started_at.map(|start| {
|
||||||
if ctx.user_id != user.user_id {
|
let end = ctx.completed_at.unwrap_or_else(chrono::Utc::now);
|
||||||
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
|
(end - start).num_seconds().max(0) as u64
|
||||||
}
|
});
|
||||||
let elapsed_secs = ctx.started_at.map(|start| {
|
|
||||||
let end = ctx.completed_at.unwrap_or_else(chrono::Utc::now);
|
|
||||||
(end - start).num_seconds().max(0) as u64
|
|
||||||
});
|
|
||||||
|
|
||||||
// Only show prompt bar for jobs that have a running worker (Pending/InProgress).
|
// Only show prompt bar for jobs that have a running worker (Pending/InProgress).
|
||||||
// Stuck jobs have no active worker loop, so messages would be silently dropped.
|
// Stuck jobs have no active worker loop, so messages would be silently dropped.
|
||||||
let is_promptable = matches!(
|
let is_promptable = matches!(
|
||||||
ctx.state,
|
ctx.state,
|
||||||
crate::context::JobState::Pending | crate::context::JobState::InProgress
|
crate::context::JobState::Pending | crate::context::JobState::InProgress
|
||||||
);
|
);
|
||||||
Ok(Json(JobDetailResponse {
|
return Ok(Json(JobDetailResponse {
|
||||||
id: ctx.job_id,
|
id: ctx.job_id,
|
||||||
title: ctx.title.clone(),
|
title: ctx.title.clone(),
|
||||||
description: ctx.description.clone(),
|
description: ctx.description.clone(),
|
||||||
state: ctx.state.to_string(),
|
state: ctx.state.to_string(),
|
||||||
user_id: ctx.user_id.clone(),
|
user_id: ctx.user_id.clone(),
|
||||||
created_at: ctx.created_at.to_rfc3339(),
|
created_at: ctx.created_at.to_rfc3339(),
|
||||||
started_at: ctx.started_at.map(|dt| dt.to_rfc3339()),
|
started_at: ctx.started_at.map(|dt| dt.to_rfc3339()),
|
||||||
completed_at: ctx.completed_at.map(|dt| dt.to_rfc3339()),
|
completed_at: ctx.completed_at.map(|dt| dt.to_rfc3339()),
|
||||||
elapsed_secs,
|
elapsed_secs,
|
||||||
project_dir: None,
|
project_dir: None,
|
||||||
browse_url: None,
|
browse_url: None,
|
||||||
job_mode: None,
|
job_mode: None,
|
||||||
transitions: Vec::new(),
|
transitions: Vec::new(),
|
||||||
can_restart: state.scheduler.is_some(),
|
can_restart: state.scheduler.is_some(),
|
||||||
can_prompt: is_promptable && state.scheduler.is_some(),
|
can_prompt: is_promptable && state.scheduler.is_some(),
|
||||||
job_kind: Some("agent".to_string()),
|
job_kind: Some("agent".to_string()),
|
||||||
}))
|
}));
|
||||||
}
|
|
||||||
Ok(None) => Err((StatusCode::NOT_FOUND, "Job not found".to_string())),
|
|
||||||
Err(e) => Err((
|
|
||||||
StatusCode::INTERNAL_SERVER_ERROR,
|
|
||||||
format!("Database error: {}", e),
|
|
||||||
)),
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
Err((StatusCode::NOT_FOUND, "Job not found".to_string()))
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn jobs_cancel_handler(
|
pub async fn jobs_cancel_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Path(id): Path<String>,
|
Path(id): Path<String>,
|
||||||
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
||||||
let job_id = Uuid::parse_str(&id)
|
let job_id = Uuid::parse_str(&id)
|
||||||
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
|
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
|
||||||
|
|
||||||
// Try sandbox job cancellation.
|
// Try sandbox job cancellation.
|
||||||
if let Some(ref store) = state.store {
|
if let Some(ref store) = state.store
|
||||||
match store.get_sandbox_job(job_id).await {
|
&& let Ok(Some(job)) = store.get_sandbox_job(job_id).await
|
||||||
Ok(Some(job)) => {
|
{
|
||||||
if job.user_id != user.user_id {
|
if job.status == "running" || job.status == "creating" {
|
||||||
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
|
// Stop the container if we have a job manager.
|
||||||
}
|
if let Some(ref jm) = state.job_manager
|
||||||
if job.status == "running" || job.status == "creating" {
|
&& let Err(e) = jm.stop_job(job_id).await
|
||||||
if let Some(ref jm) = state.job_manager
|
{
|
||||||
&& let Err(e) = jm.stop_job(job_id).await
|
tracing::warn!(job_id = %job_id, error = %e, "Failed to stop container during cancellation");
|
||||||
{
|
|
||||||
tracing::warn!(job_id = %job_id, error = %e, "Failed to stop container during cancellation");
|
|
||||||
}
|
|
||||||
store
|
|
||||||
.update_sandbox_job_status(
|
|
||||||
job_id,
|
|
||||||
"failed",
|
|
||||||
Some(false),
|
|
||||||
Some("Cancelled by user"),
|
|
||||||
None,
|
|
||||||
Some(chrono::Utc::now()),
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
|
||||||
}
|
|
||||||
return Ok(Json(serde_json::json!({
|
|
||||||
"status": "cancelled",
|
|
||||||
"job_id": job_id,
|
|
||||||
})));
|
|
||||||
}
|
|
||||||
Ok(None) => {}
|
|
||||||
Err(e) => {
|
|
||||||
return Err((
|
|
||||||
StatusCode::INTERNAL_SERVER_ERROR,
|
|
||||||
format!("Database error: {}", e),
|
|
||||||
));
|
|
||||||
}
|
}
|
||||||
|
store
|
||||||
|
.update_sandbox_job_status(
|
||||||
|
job_id,
|
||||||
|
"failed",
|
||||||
|
Some(false),
|
||||||
|
Some("Cancelled by user"),
|
||||||
|
None,
|
||||||
|
Some(chrono::Utc::now()),
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||||
}
|
}
|
||||||
|
return Ok(Json(serde_json::json!({
|
||||||
|
"status": "cancelled",
|
||||||
|
"job_id": job_id,
|
||||||
|
})));
|
||||||
}
|
}
|
||||||
|
|
||||||
// Fall back to agent job cancellation: stop the worker via the scheduler
|
// Fall back to agent job cancellation: stop the worker via the scheduler
|
||||||
// (which updates the in-memory ContextManager AND aborts the task handle),
|
// (which updates the in-memory ContextManager AND aborts the task handle),
|
||||||
// then persist the status to the DB as a fallback.
|
// then persist the status to the DB as a fallback.
|
||||||
if let Some(ref store) = state.store {
|
if let Some(ref store) = state.store
|
||||||
match store.get_job(job_id).await {
|
&& let Ok(Some(job)) = store.get_job(job_id).await
|
||||||
Ok(Some(job)) => {
|
{
|
||||||
if job.user_id != user.user_id {
|
if job.state.is_active() {
|
||||||
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
|
// Try to stop via scheduler (aborts the worker task + updates
|
||||||
}
|
// in-memory ContextManager). This is best-effort — the job may
|
||||||
if job.state.is_active() {
|
// not be in the scheduler map if it already finished.
|
||||||
// Try to stop via scheduler (aborts the worker task + updates
|
if let Some(ref slot) = state.scheduler
|
||||||
// in-memory ContextManager). This is best-effort — the job may
|
&& let Some(ref scheduler) = *slot.read().await
|
||||||
// not be in the scheduler map if it already finished.
|
{
|
||||||
if let Some(ref slot) = state.scheduler
|
let _ = scheduler.stop(job_id).await;
|
||||||
&& let Some(ref scheduler) = *slot.read().await
|
}
|
||||||
{
|
|
||||||
let _ = scheduler.stop(job_id).await;
|
|
||||||
}
|
|
||||||
|
|
||||||
// Always persist cancellation to the DB so the state is
|
// Always persist cancellation to the DB so the state is
|
||||||
// consistent even if the scheduler wasn't available or the
|
// consistent even if the scheduler wasn't available or the
|
||||||
// job wasn't in its in-memory map.
|
// job wasn't in its in-memory map.
|
||||||
store
|
store
|
||||||
.update_job_status(
|
.update_job_status(
|
||||||
job_id,
|
job_id,
|
||||||
crate::context::JobState::Cancelled,
|
crate::context::JobState::Cancelled,
|
||||||
Some("Cancelled by user"),
|
Some("Cancelled by user"),
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||||
}
|
|
||||||
return Ok(Json(serde_json::json!({
|
|
||||||
"status": "cancelled",
|
|
||||||
"job_id": job_id,
|
|
||||||
})));
|
|
||||||
}
|
|
||||||
Ok(None) => {}
|
|
||||||
Err(e) => {
|
|
||||||
return Err((
|
|
||||||
StatusCode::INTERNAL_SERVER_ERROR,
|
|
||||||
format!("Database error: {}", e),
|
|
||||||
));
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
return Ok(Json(serde_json::json!({
|
||||||
|
"status": "cancelled",
|
||||||
|
"job_id": job_id,
|
||||||
|
})));
|
||||||
}
|
}
|
||||||
|
|
||||||
Err((StatusCode::NOT_FOUND, "Job not found".to_string()))
|
Err((StatusCode::NOT_FOUND, "Job not found".to_string()))
|
||||||
@@ -363,7 +315,6 @@ pub async fn jobs_cancel_handler(
|
|||||||
|
|
||||||
pub async fn jobs_restart_handler(
|
pub async fn jobs_restart_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Path(id): Path<String>,
|
Path(id): Path<String>,
|
||||||
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
||||||
let store = state.store.as_ref().ok_or((
|
let store = state.store.as_ref().ok_or((
|
||||||
@@ -375,166 +326,146 @@ pub async fn jobs_restart_handler(
|
|||||||
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
|
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
|
||||||
|
|
||||||
// Try sandbox job restart first.
|
// Try sandbox job restart first.
|
||||||
match store.get_sandbox_job(old_job_id).await {
|
if let Ok(Some(old_job)) = store.get_sandbox_job(old_job_id).await {
|
||||||
Ok(Some(old_job)) => {
|
if old_job.status != "interrupted" && old_job.status != "failed" {
|
||||||
if old_job.user_id != user.user_id {
|
|
||||||
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
|
|
||||||
}
|
|
||||||
if old_job.status != "interrupted" && old_job.status != "failed" {
|
|
||||||
return Err((
|
|
||||||
StatusCode::CONFLICT,
|
|
||||||
format!("Cannot restart job in state '{}'", old_job.status),
|
|
||||||
));
|
|
||||||
}
|
|
||||||
|
|
||||||
let jm = state.job_manager.as_ref().ok_or((
|
|
||||||
StatusCode::SERVICE_UNAVAILABLE,
|
|
||||||
"Sandbox not enabled".to_string(),
|
|
||||||
))?;
|
|
||||||
|
|
||||||
// Enrich the task with failure context.
|
|
||||||
let task = if let Some(ref reason) = old_job.failure_reason {
|
|
||||||
format!(
|
|
||||||
"Previous attempt failed: {}. Retry: {}",
|
|
||||||
reason, old_job.task
|
|
||||||
)
|
|
||||||
} else {
|
|
||||||
old_job.task.clone()
|
|
||||||
};
|
|
||||||
|
|
||||||
let new_job_id = Uuid::new_v4();
|
|
||||||
let now = chrono::Utc::now();
|
|
||||||
|
|
||||||
let record = crate::history::SandboxJobRecord {
|
|
||||||
id: new_job_id,
|
|
||||||
task: task.clone(),
|
|
||||||
status: "creating".to_string(),
|
|
||||||
user_id: old_job.user_id.clone(),
|
|
||||||
project_dir: old_job.project_dir.clone(),
|
|
||||||
success: None,
|
|
||||||
failure_reason: None,
|
|
||||||
created_at: now,
|
|
||||||
started_at: None,
|
|
||||||
completed_at: None,
|
|
||||||
credential_grants_json: old_job.credential_grants_json.clone(),
|
|
||||||
};
|
|
||||||
store
|
|
||||||
.save_sandbox_job(&record)
|
|
||||||
.await
|
|
||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
|
||||||
|
|
||||||
let mode = match store.get_sandbox_job_mode(old_job_id).await {
|
|
||||||
Ok(Some(m)) if m == "claude_code" => {
|
|
||||||
crate::orchestrator::job_manager::JobMode::ClaudeCode
|
|
||||||
}
|
|
||||||
_ => crate::orchestrator::job_manager::JobMode::Worker,
|
|
||||||
};
|
|
||||||
|
|
||||||
let credential_grants: Vec<crate::orchestrator::auth::CredentialGrant> =
|
|
||||||
serde_json::from_str(&old_job.credential_grants_json).unwrap_or_else(|e| {
|
|
||||||
tracing::warn!(
|
|
||||||
job_id = %old_job.id,
|
|
||||||
"Failed to deserialize credential grants from stored job: {}. \
|
|
||||||
Restarted job will have no credentials.",
|
|
||||||
e
|
|
||||||
);
|
|
||||||
vec![]
|
|
||||||
});
|
|
||||||
|
|
||||||
let project_dir = std::path::PathBuf::from(&old_job.project_dir);
|
|
||||||
let _token = jm
|
|
||||||
.create_job(
|
|
||||||
new_job_id,
|
|
||||||
&task,
|
|
||||||
Some(project_dir),
|
|
||||||
mode,
|
|
||||||
credential_grants,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(|e| {
|
|
||||||
(
|
|
||||||
StatusCode::INTERNAL_SERVER_ERROR,
|
|
||||||
format!("Failed to create container: {}", e),
|
|
||||||
)
|
|
||||||
})?;
|
|
||||||
|
|
||||||
store
|
|
||||||
.update_sandbox_job_status(new_job_id, "running", None, None, Some(now), None)
|
|
||||||
.await
|
|
||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
|
||||||
|
|
||||||
return Ok(Json(serde_json::json!({
|
|
||||||
"status": "restarted",
|
|
||||||
"old_job_id": old_job_id,
|
|
||||||
"new_job_id": new_job_id,
|
|
||||||
})));
|
|
||||||
}
|
|
||||||
Ok(None) => {}
|
|
||||||
Err(e) => {
|
|
||||||
return Err((
|
return Err((
|
||||||
StatusCode::INTERNAL_SERVER_ERROR,
|
StatusCode::CONFLICT,
|
||||||
format!("Database error: {}", e),
|
format!("Cannot restart job in state '{}'", old_job.status),
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
let jm = state.job_manager.as_ref().ok_or((
|
||||||
|
StatusCode::SERVICE_UNAVAILABLE,
|
||||||
|
"Sandbox not enabled".to_string(),
|
||||||
|
))?;
|
||||||
|
|
||||||
|
// Enrich the task with failure context.
|
||||||
|
let task = if let Some(ref reason) = old_job.failure_reason {
|
||||||
|
format!(
|
||||||
|
"Previous attempt failed: {}. Retry: {}",
|
||||||
|
reason, old_job.task
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
old_job.task.clone()
|
||||||
|
};
|
||||||
|
|
||||||
|
let new_job_id = Uuid::new_v4();
|
||||||
|
let now = chrono::Utc::now();
|
||||||
|
|
||||||
|
let record = crate::history::SandboxJobRecord {
|
||||||
|
id: new_job_id,
|
||||||
|
task: task.clone(),
|
||||||
|
status: "creating".to_string(),
|
||||||
|
user_id: old_job.user_id.clone(),
|
||||||
|
project_dir: old_job.project_dir.clone(),
|
||||||
|
success: None,
|
||||||
|
failure_reason: None,
|
||||||
|
created_at: now,
|
||||||
|
started_at: None,
|
||||||
|
completed_at: None,
|
||||||
|
credential_grants_json: old_job.credential_grants_json.clone(),
|
||||||
|
};
|
||||||
|
store
|
||||||
|
.save_sandbox_job(&record)
|
||||||
|
.await
|
||||||
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||||
|
|
||||||
|
let mode = match store.get_sandbox_job_mode(old_job_id).await {
|
||||||
|
Ok(Some(m)) if m == "claude_code" => {
|
||||||
|
crate::orchestrator::job_manager::JobMode::ClaudeCode
|
||||||
|
}
|
||||||
|
_ => crate::orchestrator::job_manager::JobMode::Worker,
|
||||||
|
};
|
||||||
|
|
||||||
|
let credential_grants: Vec<crate::orchestrator::auth::CredentialGrant> =
|
||||||
|
serde_json::from_str(&old_job.credential_grants_json).unwrap_or_else(|e| {
|
||||||
|
tracing::warn!(
|
||||||
|
job_id = %old_job.id,
|
||||||
|
"Failed to deserialize credential grants from stored job: {}. \
|
||||||
|
Restarted job will have no credentials.",
|
||||||
|
e
|
||||||
|
);
|
||||||
|
vec![]
|
||||||
|
});
|
||||||
|
|
||||||
|
let project_dir = std::path::PathBuf::from(&old_job.project_dir);
|
||||||
|
let _token = jm
|
||||||
|
.create_job(
|
||||||
|
new_job_id,
|
||||||
|
&task,
|
||||||
|
Some(project_dir),
|
||||||
|
mode,
|
||||||
|
credential_grants,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.map_err(|e| {
|
||||||
|
(
|
||||||
|
StatusCode::INTERNAL_SERVER_ERROR,
|
||||||
|
format!("Failed to create container: {}", e),
|
||||||
|
)
|
||||||
|
})?;
|
||||||
|
|
||||||
|
store
|
||||||
|
.update_sandbox_job_status(new_job_id, "running", None, None, Some(now), None)
|
||||||
|
.await
|
||||||
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||||
|
|
||||||
|
return Ok(Json(serde_json::json!({
|
||||||
|
"status": "restarted",
|
||||||
|
"old_job_id": old_job_id,
|
||||||
|
"new_job_id": new_job_id,
|
||||||
|
})));
|
||||||
}
|
}
|
||||||
|
|
||||||
// Try agent job restart: dispatch a new job via the scheduler.
|
// Try agent job restart: dispatch a new job via the scheduler.
|
||||||
match store.get_job(old_job_id).await {
|
if let Ok(Some(old_job)) = store.get_job(old_job_id).await {
|
||||||
Ok(Some(old_job)) => {
|
if old_job.state.is_active() {
|
||||||
if old_job.user_id != user.user_id {
|
return Err((
|
||||||
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
|
StatusCode::CONFLICT,
|
||||||
}
|
format!("Cannot restart job in state '{}'", old_job.state),
|
||||||
if old_job.state.is_active() {
|
));
|
||||||
return Err((
|
|
||||||
StatusCode::CONFLICT,
|
|
||||||
format!("Cannot restart job in state '{}'", old_job.state),
|
|
||||||
));
|
|
||||||
}
|
|
||||||
|
|
||||||
let slot = state.scheduler.as_ref().ok_or((
|
|
||||||
StatusCode::SERVICE_UNAVAILABLE,
|
|
||||||
"Scheduler not available".to_string(),
|
|
||||||
))?;
|
|
||||||
let scheduler_guard = slot.read().await;
|
|
||||||
let scheduler = scheduler_guard.as_ref().ok_or((
|
|
||||||
StatusCode::SERVICE_UNAVAILABLE,
|
|
||||||
"Agent not started yet".to_string(),
|
|
||||||
))?;
|
|
||||||
|
|
||||||
// Look up failure reason (O(1) point lookup).
|
|
||||||
let failure_reason = store
|
|
||||||
.get_agent_job_failure_reason(old_job_id)
|
|
||||||
.await
|
|
||||||
.ok()
|
|
||||||
.flatten()
|
|
||||||
.unwrap_or_default();
|
|
||||||
|
|
||||||
let title = if !failure_reason.is_empty() {
|
|
||||||
format!(
|
|
||||||
"Previous attempt failed: {}. Retry: {}",
|
|
||||||
failure_reason, old_job.title
|
|
||||||
)
|
|
||||||
} else {
|
|
||||||
old_job.title.clone()
|
|
||||||
};
|
|
||||||
|
|
||||||
let new_job_id = scheduler
|
|
||||||
.dispatch_job(&old_job.user_id, &title, &old_job.description, None)
|
|
||||||
.await
|
|
||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
|
||||||
|
|
||||||
Ok(Json(serde_json::json!({
|
|
||||||
"status": "restarted",
|
|
||||||
"old_job_id": old_job_id,
|
|
||||||
"new_job_id": new_job_id,
|
|
||||||
})))
|
|
||||||
}
|
}
|
||||||
Ok(None) => Err((StatusCode::NOT_FOUND, "Job not found".to_string())),
|
|
||||||
Err(e) => Err((
|
let slot = state.scheduler.as_ref().ok_or((
|
||||||
StatusCode::INTERNAL_SERVER_ERROR,
|
StatusCode::SERVICE_UNAVAILABLE,
|
||||||
format!("Database error: {}", e),
|
"Scheduler not available".to_string(),
|
||||||
)),
|
))?;
|
||||||
|
let scheduler_guard = slot.read().await;
|
||||||
|
let scheduler = scheduler_guard.as_ref().ok_or((
|
||||||
|
StatusCode::SERVICE_UNAVAILABLE,
|
||||||
|
"Agent not started yet".to_string(),
|
||||||
|
))?;
|
||||||
|
|
||||||
|
// Look up failure reason (O(1) point lookup).
|
||||||
|
let failure_reason = store
|
||||||
|
.get_agent_job_failure_reason(old_job_id)
|
||||||
|
.await
|
||||||
|
.ok()
|
||||||
|
.flatten()
|
||||||
|
.unwrap_or_default();
|
||||||
|
|
||||||
|
let title = if !failure_reason.is_empty() {
|
||||||
|
format!(
|
||||||
|
"Previous attempt failed: {}. Retry: {}",
|
||||||
|
failure_reason, old_job.title
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
old_job.title.clone()
|
||||||
|
};
|
||||||
|
|
||||||
|
let new_job_id = scheduler
|
||||||
|
.dispatch_job(&old_job.user_id, &title, &old_job.description, None)
|
||||||
|
.await
|
||||||
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||||
|
|
||||||
|
return Ok(Json(serde_json::json!({
|
||||||
|
"status": "restarted",
|
||||||
|
"old_job_id": old_job_id,
|
||||||
|
"new_job_id": new_job_id,
|
||||||
|
})));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
Err((StatusCode::NOT_FOUND, "Job not found".to_string()))
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Submit a follow-up prompt to a running job.
|
/// Submit a follow-up prompt to a running job.
|
||||||
@@ -545,7 +476,6 @@ pub async fn jobs_restart_handler(
|
|||||||
/// - Worker-mode sandbox jobs → not supported (no mechanism to inject)
|
/// - Worker-mode sandbox jobs → not supported (no mechanism to inject)
|
||||||
pub async fn jobs_prompt_handler(
|
pub async fn jobs_prompt_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Path(id): Path<String>,
|
Path(id): Path<String>,
|
||||||
Json(body): Json<serde_json::Value>,
|
Json(body): Json<serde_json::Value>,
|
||||||
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
||||||
@@ -564,15 +494,10 @@ pub async fn jobs_prompt_handler(
|
|||||||
|
|
||||||
let done = body.get("done").and_then(|v| v.as_bool()).unwrap_or(false);
|
let done = body.get("done").and_then(|v| v.as_bool()).unwrap_or(false);
|
||||||
|
|
||||||
// Try sandbox job path first: verify ownership, then route to Claude Code or reject.
|
// Try sandbox job path: check if we have a sandbox record for this ID.
|
||||||
if let Some(ref s) = state.store
|
if let Some(ref s) = state.store
|
||||||
&& let Ok(Some(sandbox_job)) = s.get_sandbox_job(job_id).await
|
&& let Ok(Some(_)) = s.get_sandbox_job(job_id).await
|
||||||
{
|
{
|
||||||
// Verify ownership.
|
|
||||||
if sandbox_job.user_id != user.user_id {
|
|
||||||
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
|
|
||||||
}
|
|
||||||
|
|
||||||
// It's a sandbox job. Check if Claude Code mode.
|
// It's a sandbox job. Check if Claude Code mode.
|
||||||
let mode = s.get_sandbox_job_mode(job_id).await.ok().flatten();
|
let mode = s.get_sandbox_job_mode(job_id).await.ok().flatten();
|
||||||
if mode.as_deref() == Some("claude_code") {
|
if mode.as_deref() == Some("claude_code") {
|
||||||
@@ -597,14 +522,7 @@ pub async fn jobs_prompt_handler(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Try agent job path: verify ownership, then send via scheduler.
|
// Try agent job path: send via scheduler.
|
||||||
if let Some(ref store) = state.store
|
|
||||||
&& let Ok(Some(agent_job)) = store.get_job(job_id).await
|
|
||||||
&& agent_job.user_id != user.user_id
|
|
||||||
{
|
|
||||||
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
|
|
||||||
}
|
|
||||||
|
|
||||||
let slot = state.scheduler.as_ref().ok_or((
|
let slot = state.scheduler.as_ref().ok_or((
|
||||||
StatusCode::NOT_IMPLEMENTED,
|
StatusCode::NOT_IMPLEMENTED,
|
||||||
"Agent job prompts require the scheduler to be configured".to_string(),
|
"Agent job prompts require the scheduler to be configured".to_string(),
|
||||||
@@ -632,7 +550,6 @@ pub async fn jobs_prompt_handler(
|
|||||||
/// Load persisted job events for a job (for history replay on page open).
|
/// Load persisted job events for a job (for history replay on page open).
|
||||||
pub async fn jobs_events_handler(
|
pub async fn jobs_events_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Path(id): Path<String>,
|
Path(id): Path<String>,
|
||||||
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
||||||
let store = state.store.as_ref().ok_or((
|
let store = state.store.as_ref().ok_or((
|
||||||
@@ -644,24 +561,6 @@ pub async fn jobs_events_handler(
|
|||||||
.parse()
|
.parse()
|
||||||
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
|
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
|
||||||
|
|
||||||
// Verify ownership before returning events.
|
|
||||||
match store.get_sandbox_job(job_id).await {
|
|
||||||
Ok(Some(job)) => {
|
|
||||||
if job.user_id != user.user_id {
|
|
||||||
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
Ok(None) => {
|
|
||||||
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
|
|
||||||
}
|
|
||||||
Err(e) => {
|
|
||||||
return Err((
|
|
||||||
StatusCode::INTERNAL_SERVER_ERROR,
|
|
||||||
format!("Database error: {}", e),
|
|
||||||
));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
let events = store
|
let events = store
|
||||||
.list_job_events(job_id, None)
|
.list_job_events(job_id, None)
|
||||||
.await
|
.await
|
||||||
@@ -694,7 +593,6 @@ pub struct FilePathQuery {
|
|||||||
|
|
||||||
pub async fn job_files_list_handler(
|
pub async fn job_files_list_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Path(id): Path<String>,
|
Path(id): Path<String>,
|
||||||
Query(query): Query<FilePathQuery>,
|
Query(query): Query<FilePathQuery>,
|
||||||
) -> Result<Json<ProjectFilesResponse>, (StatusCode, String)> {
|
) -> Result<Json<ProjectFilesResponse>, (StatusCode, String)> {
|
||||||
@@ -712,10 +610,6 @@ pub async fn job_files_list_handler(
|
|||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
||||||
.ok_or((StatusCode::NOT_FOUND, "Job not found".to_string()))?;
|
.ok_or((StatusCode::NOT_FOUND, "Job not found".to_string()))?;
|
||||||
|
|
||||||
if job.user_id != user.user_id {
|
|
||||||
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
|
|
||||||
}
|
|
||||||
|
|
||||||
let base = std::path::PathBuf::from(&job.project_dir);
|
let base = std::path::PathBuf::from(&job.project_dir);
|
||||||
let rel_path = query.path.as_deref().unwrap_or("");
|
let rel_path = query.path.as_deref().unwrap_or("");
|
||||||
let target = base.join(rel_path);
|
let target = base.join(rel_path);
|
||||||
@@ -762,7 +656,6 @@ pub async fn job_files_list_handler(
|
|||||||
|
|
||||||
pub async fn job_files_read_handler(
|
pub async fn job_files_read_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Path(id): Path<String>,
|
Path(id): Path<String>,
|
||||||
Query(query): Query<FilePathQuery>,
|
Query(query): Query<FilePathQuery>,
|
||||||
) -> Result<Json<ProjectFileReadResponse>, (StatusCode, String)> {
|
) -> Result<Json<ProjectFileReadResponse>, (StatusCode, String)> {
|
||||||
@@ -780,10 +673,6 @@ pub async fn job_files_read_handler(
|
|||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
||||||
.ok_or((StatusCode::NOT_FOUND, "Job not found".to_string()))?;
|
.ok_or((StatusCode::NOT_FOUND, "Job not found".to_string()))?;
|
||||||
|
|
||||||
if job.user_id != user.user_id {
|
|
||||||
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
|
|
||||||
}
|
|
||||||
|
|
||||||
let path = query.path.as_deref().ok_or((
|
let path = query.path.as_deref().ok_or((
|
||||||
StatusCode::BAD_REQUEST,
|
StatusCode::BAD_REQUEST,
|
||||||
"path parameter required".to_string(),
|
"path parameter required".to_string(),
|
||||||
|
|||||||
@@ -0,0 +1,154 @@
|
|||||||
|
//! Memory/workspace API handlers.
|
||||||
|
|
||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
|
use axum::{
|
||||||
|
Json,
|
||||||
|
extract::{Query, State},
|
||||||
|
http::StatusCode,
|
||||||
|
};
|
||||||
|
use serde::Deserialize;
|
||||||
|
|
||||||
|
use crate::channels::web::server::GatewayState;
|
||||||
|
use crate::channels::web::types::*;
|
||||||
|
|
||||||
|
#[derive(Deserialize)]
|
||||||
|
pub struct TreeQuery {
|
||||||
|
#[allow(dead_code)]
|
||||||
|
pub depth: Option<usize>,
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn memory_tree_handler(
|
||||||
|
State(state): State<Arc<GatewayState>>,
|
||||||
|
Query(_query): Query<TreeQuery>,
|
||||||
|
) -> Result<Json<MemoryTreeResponse>, (StatusCode, String)> {
|
||||||
|
let workspace = state.workspace.as_ref().ok_or((
|
||||||
|
StatusCode::SERVICE_UNAVAILABLE,
|
||||||
|
"Workspace not available".to_string(),
|
||||||
|
))?;
|
||||||
|
|
||||||
|
// Build tree from list_all (flat list of all paths)
|
||||||
|
let all_paths = workspace
|
||||||
|
.list_all()
|
||||||
|
.await
|
||||||
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||||
|
|
||||||
|
// Collect unique directories and files
|
||||||
|
let mut entries: Vec<TreeEntry> = Vec::new();
|
||||||
|
let mut seen_dirs: std::collections::HashSet<String> = std::collections::HashSet::new();
|
||||||
|
|
||||||
|
for path in &all_paths {
|
||||||
|
// Add parent directories
|
||||||
|
let parts: Vec<&str> = path.split('/').collect();
|
||||||
|
for i in 0..parts.len().saturating_sub(1) {
|
||||||
|
let dir_path = parts[..=i].join("/");
|
||||||
|
if seen_dirs.insert(dir_path.clone()) {
|
||||||
|
entries.push(TreeEntry {
|
||||||
|
path: dir_path,
|
||||||
|
is_dir: true,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Add the file itself
|
||||||
|
entries.push(TreeEntry {
|
||||||
|
path: path.clone(),
|
||||||
|
is_dir: false,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
entries.sort_by(|a, b| a.path.cmp(&b.path));
|
||||||
|
|
||||||
|
Ok(Json(MemoryTreeResponse { entries }))
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Deserialize)]
|
||||||
|
pub struct ListQuery {
|
||||||
|
pub path: Option<String>,
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn memory_list_handler(
|
||||||
|
State(state): State<Arc<GatewayState>>,
|
||||||
|
Query(query): Query<ListQuery>,
|
||||||
|
) -> Result<Json<MemoryListResponse>, (StatusCode, String)> {
|
||||||
|
let workspace = state.workspace.as_ref().ok_or((
|
||||||
|
StatusCode::SERVICE_UNAVAILABLE,
|
||||||
|
"Workspace not available".to_string(),
|
||||||
|
))?;
|
||||||
|
|
||||||
|
let path = query.path.as_deref().unwrap_or("");
|
||||||
|
let entries = workspace
|
||||||
|
.list(path)
|
||||||
|
.await
|
||||||
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||||
|
|
||||||
|
let list_entries: Vec<ListEntry> = entries
|
||||||
|
.iter()
|
||||||
|
.map(|e| ListEntry {
|
||||||
|
name: e.path.rsplit('/').next().unwrap_or(&e.path).to_string(),
|
||||||
|
path: e.path.clone(),
|
||||||
|
is_dir: e.is_directory,
|
||||||
|
updated_at: e.updated_at.map(|dt| dt.to_rfc3339()),
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
|
||||||
|
Ok(Json(MemoryListResponse {
|
||||||
|
path: path.to_string(),
|
||||||
|
entries: list_entries,
|
||||||
|
}))
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Deserialize)]
|
||||||
|
pub struct ReadQuery {
|
||||||
|
pub path: String,
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn memory_read_handler(
|
||||||
|
State(state): State<Arc<GatewayState>>,
|
||||||
|
Query(query): Query<ReadQuery>,
|
||||||
|
) -> Result<Json<MemoryReadResponse>, (StatusCode, String)> {
|
||||||
|
let workspace = state.workspace.as_ref().ok_or((
|
||||||
|
StatusCode::SERVICE_UNAVAILABLE,
|
||||||
|
"Workspace not available".to_string(),
|
||||||
|
))?;
|
||||||
|
|
||||||
|
let doc = workspace
|
||||||
|
.read(&query.path)
|
||||||
|
.await
|
||||||
|
.map_err(|e| (StatusCode::NOT_FOUND, e.to_string()))?;
|
||||||
|
|
||||||
|
Ok(Json(MemoryReadResponse {
|
||||||
|
path: query.path,
|
||||||
|
content: doc.content,
|
||||||
|
updated_at: Some(doc.updated_at.to_rfc3339()),
|
||||||
|
}))
|
||||||
|
}
|
||||||
|
|
||||||
|
// memory_write_handler lives in server.rs (layer-aware version with append,
|
||||||
|
// privacy redirect, and proper error status codes).
|
||||||
|
|
||||||
|
pub async fn memory_search_handler(
|
||||||
|
State(state): State<Arc<GatewayState>>,
|
||||||
|
Json(req): Json<MemorySearchRequest>,
|
||||||
|
) -> Result<Json<MemorySearchResponse>, (StatusCode, String)> {
|
||||||
|
let workspace = state.workspace.as_ref().ok_or((
|
||||||
|
StatusCode::SERVICE_UNAVAILABLE,
|
||||||
|
"Workspace not available".to_string(),
|
||||||
|
))?;
|
||||||
|
|
||||||
|
let limit = req.limit.unwrap_or(10);
|
||||||
|
let results = workspace
|
||||||
|
.search(&req.query, limit)
|
||||||
|
.await
|
||||||
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||||
|
|
||||||
|
let hits: Vec<SearchHit> = results
|
||||||
|
.into_iter()
|
||||||
|
.map(|r| SearchHit {
|
||||||
|
path: r.document_path,
|
||||||
|
content: r.content,
|
||||||
|
score: r.score as f64,
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
|
||||||
|
Ok(Json(MemorySearchResponse { results: hits }))
|
||||||
|
}
|
||||||
@@ -1,9 +1,13 @@
|
|||||||
//! Handler modules for the web gateway API.
|
//! Handler modules for the web gateway API.
|
||||||
//!
|
//!
|
||||||
//! Each module groups related endpoint handlers by domain.
|
//! Each module groups related endpoint handlers by domain.
|
||||||
|
//!
|
||||||
|
//! # Migration status
|
||||||
|
//!
|
||||||
|
//! `skills` is the canonical implementation used by `server.rs`.
|
||||||
|
//! The remaining modules are in-progress migrations from inline server.rs
|
||||||
|
//! handlers; their functions are not yet wired up, hence the `dead_code` allow.
|
||||||
|
|
||||||
pub mod jobs;
|
|
||||||
pub mod routines;
|
|
||||||
pub mod skills;
|
pub mod skills;
|
||||||
|
|
||||||
// Modules not yet wired into server.rs router -- suppress dead_code until
|
// Modules not yet wired into server.rs router -- suppress dead_code until
|
||||||
@@ -13,6 +17,12 @@ pub mod chat;
|
|||||||
#[allow(dead_code)]
|
#[allow(dead_code)]
|
||||||
pub mod extensions;
|
pub mod extensions;
|
||||||
#[allow(dead_code)]
|
#[allow(dead_code)]
|
||||||
|
pub mod jobs;
|
||||||
|
#[allow(dead_code)]
|
||||||
|
pub mod memory;
|
||||||
|
#[allow(dead_code)]
|
||||||
|
pub mod routines;
|
||||||
|
#[allow(dead_code)]
|
||||||
pub mod settings;
|
pub mod settings;
|
||||||
#[allow(dead_code)]
|
#[allow(dead_code)]
|
||||||
pub mod static_files;
|
pub mod static_files;
|
||||||
|
|||||||
@@ -11,14 +11,12 @@ use serde::Deserialize;
|
|||||||
use uuid::Uuid;
|
use uuid::Uuid;
|
||||||
|
|
||||||
use crate::agent::routine::{Trigger, next_cron_fire};
|
use crate::agent::routine::{Trigger, next_cron_fire};
|
||||||
use crate::channels::web::auth::AuthenticatedUser;
|
|
||||||
use crate::channels::web::server::GatewayState;
|
use crate::channels::web::server::GatewayState;
|
||||||
use crate::channels::web::types::*;
|
use crate::channels::web::types::*;
|
||||||
use crate::error::RoutineError;
|
use crate::error::RoutineError;
|
||||||
|
|
||||||
pub async fn routines_list_handler(
|
pub async fn routines_list_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
) -> Result<Json<RoutineListResponse>, (StatusCode, String)> {
|
) -> Result<Json<RoutineListResponse>, (StatusCode, String)> {
|
||||||
let store = state.store.as_ref().ok_or((
|
let store = state.store.as_ref().ok_or((
|
||||||
StatusCode::SERVICE_UNAVAILABLE,
|
StatusCode::SERVICE_UNAVAILABLE,
|
||||||
@@ -26,7 +24,7 @@ pub async fn routines_list_handler(
|
|||||||
))?;
|
))?;
|
||||||
|
|
||||||
let routines = store
|
let routines = store
|
||||||
.list_routines(&user.user_id)
|
.list_all_routines()
|
||||||
.await
|
.await
|
||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||||
|
|
||||||
@@ -37,7 +35,6 @@ pub async fn routines_list_handler(
|
|||||||
|
|
||||||
pub async fn routines_summary_handler(
|
pub async fn routines_summary_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
) -> Result<Json<RoutineSummaryResponse>, (StatusCode, String)> {
|
) -> Result<Json<RoutineSummaryResponse>, (StatusCode, String)> {
|
||||||
let store = state.store.as_ref().ok_or((
|
let store = state.store.as_ref().ok_or((
|
||||||
StatusCode::SERVICE_UNAVAILABLE,
|
StatusCode::SERVICE_UNAVAILABLE,
|
||||||
@@ -45,7 +42,7 @@ pub async fn routines_summary_handler(
|
|||||||
))?;
|
))?;
|
||||||
|
|
||||||
let routines = store
|
let routines = store
|
||||||
.list_routines(&user.user_id)
|
.list_all_routines()
|
||||||
.await
|
.await
|
||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||||
|
|
||||||
@@ -81,7 +78,6 @@ pub async fn routines_summary_handler(
|
|||||||
|
|
||||||
pub async fn routines_detail_handler(
|
pub async fn routines_detail_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Path(id): Path<String>,
|
Path(id): Path<String>,
|
||||||
) -> Result<Json<RoutineDetailResponse>, (StatusCode, String)> {
|
) -> Result<Json<RoutineDetailResponse>, (StatusCode, String)> {
|
||||||
let store = state.store.as_ref().ok_or((
|
let store = state.store.as_ref().ok_or((
|
||||||
@@ -98,10 +94,6 @@ pub async fn routines_detail_handler(
|
|||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
||||||
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
|
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
|
||||||
|
|
||||||
if routine.user_id != user.user_id {
|
|
||||||
return Err((StatusCode::NOT_FOUND, "Routine not found".to_string()));
|
|
||||||
}
|
|
||||||
|
|
||||||
let runs = store
|
let runs = store
|
||||||
.list_routine_runs(routine_id, 20)
|
.list_routine_runs(routine_id, 20)
|
||||||
.await
|
.await
|
||||||
@@ -145,7 +137,6 @@ pub async fn routines_detail_handler(
|
|||||||
|
|
||||||
pub async fn routines_trigger_handler(
|
pub async fn routines_trigger_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Path(id): Path<String>,
|
Path(id): Path<String>,
|
||||||
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
||||||
// Clone the Arc out of the lock to avoid holding the RwLock across .await.
|
// Clone the Arc out of the lock to avoid holding the RwLock across .await.
|
||||||
@@ -161,7 +152,7 @@ pub async fn routines_trigger_handler(
|
|||||||
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
|
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
|
||||||
|
|
||||||
let run_id = engine
|
let run_id = engine
|
||||||
.fire_manual(routine_id, Some(&user.user_id))
|
.fire_manual(routine_id, Some(&state.user_id))
|
||||||
.await
|
.await
|
||||||
.map_err(|e| (routine_error_status(&e), e.to_string()))?;
|
.map_err(|e| (routine_error_status(&e), e.to_string()))?;
|
||||||
|
|
||||||
@@ -179,7 +170,6 @@ pub struct ToggleRequest {
|
|||||||
|
|
||||||
pub async fn routines_toggle_handler(
|
pub async fn routines_toggle_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Path(id): Path<String>,
|
Path(id): Path<String>,
|
||||||
body: Option<Json<ToggleRequest>>,
|
body: Option<Json<ToggleRequest>>,
|
||||||
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
||||||
@@ -197,10 +187,6 @@ pub async fn routines_toggle_handler(
|
|||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
||||||
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
|
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
|
||||||
|
|
||||||
if routine.user_id != user.user_id {
|
|
||||||
return Err((StatusCode::NOT_FOUND, "Routine not found".to_string()));
|
|
||||||
}
|
|
||||||
|
|
||||||
let was_enabled = routine.enabled;
|
let was_enabled = routine.enabled;
|
||||||
// If a specific value was provided, use it; otherwise toggle.
|
// If a specific value was provided, use it; otherwise toggle.
|
||||||
routine.enabled = match body {
|
routine.enabled = match body {
|
||||||
@@ -244,7 +230,6 @@ pub async fn routines_toggle_handler(
|
|||||||
|
|
||||||
pub async fn routines_delete_handler(
|
pub async fn routines_delete_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Path(id): Path<String>,
|
Path(id): Path<String>,
|
||||||
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
||||||
let store = state.store.as_ref().ok_or((
|
let store = state.store.as_ref().ok_or((
|
||||||
@@ -255,17 +240,6 @@ pub async fn routines_delete_handler(
|
|||||||
let routine_id = Uuid::parse_str(&id)
|
let routine_id = Uuid::parse_str(&id)
|
||||||
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
|
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
|
||||||
|
|
||||||
// Verify ownership before deleting.
|
|
||||||
let routine = store
|
|
||||||
.get_routine(routine_id)
|
|
||||||
.await
|
|
||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
|
||||||
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
|
|
||||||
|
|
||||||
if routine.user_id != user.user_id {
|
|
||||||
return Err((StatusCode::NOT_FOUND, "Routine not found".to_string()));
|
|
||||||
}
|
|
||||||
|
|
||||||
let deleted = store
|
let deleted = store
|
||||||
.delete_routine(routine_id)
|
.delete_routine(routine_id)
|
||||||
.await
|
.await
|
||||||
@@ -287,10 +261,8 @@ pub async fn routines_delete_handler(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[allow(dead_code)] // Used by server.rs inline version; kept in sync here for future migration.
|
|
||||||
pub async fn routines_runs_handler(
|
pub async fn routines_runs_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Path(id): Path<String>,
|
Path(id): Path<String>,
|
||||||
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
||||||
let store = state.store.as_ref().ok_or((
|
let store = state.store.as_ref().ok_or((
|
||||||
@@ -301,17 +273,6 @@ pub async fn routines_runs_handler(
|
|||||||
let routine_id = Uuid::parse_str(&id)
|
let routine_id = Uuid::parse_str(&id)
|
||||||
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
|
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
|
||||||
|
|
||||||
// Verify ownership before listing runs.
|
|
||||||
let routine = store
|
|
||||||
.get_routine(routine_id)
|
|
||||||
.await
|
|
||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
|
||||||
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
|
|
||||||
|
|
||||||
if routine.user_id != user.user_id {
|
|
||||||
return Err((StatusCode::NOT_FOUND, "Routine not found".to_string()));
|
|
||||||
}
|
|
||||||
|
|
||||||
let runs = store
|
let runs = store
|
||||||
.list_routine_runs(routine_id, 50)
|
.list_routine_runs(routine_id, 50)
|
||||||
.await
|
.await
|
||||||
|
|||||||
@@ -8,19 +8,17 @@ use axum::{
|
|||||||
http::StatusCode,
|
http::StatusCode,
|
||||||
};
|
};
|
||||||
|
|
||||||
use crate::channels::web::auth::AuthenticatedUser;
|
|
||||||
use crate::channels::web::server::GatewayState;
|
use crate::channels::web::server::GatewayState;
|
||||||
use crate::channels::web::types::*;
|
use crate::channels::web::types::*;
|
||||||
|
|
||||||
pub async fn settings_list_handler(
|
pub async fn settings_list_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
) -> Result<Json<SettingsListResponse>, StatusCode> {
|
) -> Result<Json<SettingsListResponse>, StatusCode> {
|
||||||
let store = state
|
let store = state
|
||||||
.store
|
.store
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
|
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
|
||||||
let rows = store.list_settings(&user.user_id).await.map_err(|e| {
|
let rows = store.list_settings(&state.user_id).await.map_err(|e| {
|
||||||
tracing::error!("Failed to list settings: {}", e);
|
tracing::error!("Failed to list settings: {}", e);
|
||||||
StatusCode::INTERNAL_SERVER_ERROR
|
StatusCode::INTERNAL_SERVER_ERROR
|
||||||
})?;
|
})?;
|
||||||
@@ -39,7 +37,6 @@ pub async fn settings_list_handler(
|
|||||||
|
|
||||||
pub async fn settings_get_handler(
|
pub async fn settings_get_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Path(key): Path<String>,
|
Path(key): Path<String>,
|
||||||
) -> Result<Json<SettingResponse>, StatusCode> {
|
) -> Result<Json<SettingResponse>, StatusCode> {
|
||||||
let store = state
|
let store = state
|
||||||
@@ -47,7 +44,7 @@ pub async fn settings_get_handler(
|
|||||||
.as_ref()
|
.as_ref()
|
||||||
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
|
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
|
||||||
let row = store
|
let row = store
|
||||||
.get_setting_full(&user.user_id, &key)
|
.get_setting_full(&state.user_id, &key)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| {
|
.map_err(|e| {
|
||||||
tracing::error!("Failed to get setting '{}': {}", key, e);
|
tracing::error!("Failed to get setting '{}': {}", key, e);
|
||||||
@@ -64,7 +61,6 @@ pub async fn settings_get_handler(
|
|||||||
|
|
||||||
pub async fn settings_set_handler(
|
pub async fn settings_set_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Path(key): Path<String>,
|
Path(key): Path<String>,
|
||||||
Json(body): Json<SettingWriteRequest>,
|
Json(body): Json<SettingWriteRequest>,
|
||||||
) -> Result<StatusCode, StatusCode> {
|
) -> Result<StatusCode, StatusCode> {
|
||||||
@@ -73,7 +69,7 @@ pub async fn settings_set_handler(
|
|||||||
.as_ref()
|
.as_ref()
|
||||||
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
|
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
|
||||||
store
|
store
|
||||||
.set_setting(&user.user_id, &key, &body.value)
|
.set_setting(&state.user_id, &key, &body.value)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| {
|
.map_err(|e| {
|
||||||
tracing::error!("Failed to set setting '{}': {}", key, e);
|
tracing::error!("Failed to set setting '{}': {}", key, e);
|
||||||
@@ -85,7 +81,6 @@ pub async fn settings_set_handler(
|
|||||||
|
|
||||||
pub async fn settings_delete_handler(
|
pub async fn settings_delete_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Path(key): Path<String>,
|
Path(key): Path<String>,
|
||||||
) -> Result<StatusCode, StatusCode> {
|
) -> Result<StatusCode, StatusCode> {
|
||||||
let store = state
|
let store = state
|
||||||
@@ -93,7 +88,7 @@ pub async fn settings_delete_handler(
|
|||||||
.as_ref()
|
.as_ref()
|
||||||
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
|
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
|
||||||
store
|
store
|
||||||
.delete_setting(&user.user_id, &key)
|
.delete_setting(&state.user_id, &key)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| {
|
.map_err(|e| {
|
||||||
tracing::error!("Failed to delete setting '{}': {}", key, e);
|
tracing::error!("Failed to delete setting '{}': {}", key, e);
|
||||||
@@ -105,13 +100,12 @@ pub async fn settings_delete_handler(
|
|||||||
|
|
||||||
pub async fn settings_export_handler(
|
pub async fn settings_export_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
) -> Result<Json<SettingsExportResponse>, StatusCode> {
|
) -> Result<Json<SettingsExportResponse>, StatusCode> {
|
||||||
let store = state
|
let store = state
|
||||||
.store
|
.store
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
|
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
|
||||||
let settings = store.get_all_settings(&user.user_id).await.map_err(|e| {
|
let settings = store.get_all_settings(&state.user_id).await.map_err(|e| {
|
||||||
tracing::error!("Failed to export settings: {}", e);
|
tracing::error!("Failed to export settings: {}", e);
|
||||||
StatusCode::INTERNAL_SERVER_ERROR
|
StatusCode::INTERNAL_SERVER_ERROR
|
||||||
})?;
|
})?;
|
||||||
@@ -121,7 +115,6 @@ pub async fn settings_export_handler(
|
|||||||
|
|
||||||
pub async fn settings_import_handler(
|
pub async fn settings_import_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
Json(body): Json<SettingsImportRequest>,
|
Json(body): Json<SettingsImportRequest>,
|
||||||
) -> Result<StatusCode, StatusCode> {
|
) -> Result<StatusCode, StatusCode> {
|
||||||
let store = state
|
let store = state
|
||||||
@@ -129,7 +122,7 @@ pub async fn settings_import_handler(
|
|||||||
.as_ref()
|
.as_ref()
|
||||||
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
|
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
|
||||||
store
|
store
|
||||||
.set_all_settings(&user.user_id, &body.settings)
|
.set_all_settings(&state.user_id, &body.settings)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| {
|
.map_err(|e| {
|
||||||
tracing::error!("Failed to import settings: {}", e);
|
tracing::error!("Failed to import settings: {}", e);
|
||||||
|
|||||||
@@ -8,13 +8,11 @@ use axum::{
|
|||||||
http::StatusCode,
|
http::StatusCode,
|
||||||
};
|
};
|
||||||
|
|
||||||
use crate::channels::web::auth::AuthenticatedUser;
|
|
||||||
use crate::channels::web::server::GatewayState;
|
use crate::channels::web::server::GatewayState;
|
||||||
use crate::channels::web::types::*;
|
use crate::channels::web::types::*;
|
||||||
|
|
||||||
pub async fn skills_list_handler(
|
pub async fn skills_list_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(_user): AuthenticatedUser,
|
|
||||||
) -> Result<Json<SkillListResponse>, (StatusCode, String)> {
|
) -> Result<Json<SkillListResponse>, (StatusCode, String)> {
|
||||||
let registry = state.skill_registry.as_ref().ok_or((
|
let registry = state.skill_registry.as_ref().ok_or((
|
||||||
StatusCode::NOT_IMPLEMENTED,
|
StatusCode::NOT_IMPLEMENTED,
|
||||||
@@ -47,7 +45,6 @@ pub async fn skills_list_handler(
|
|||||||
|
|
||||||
pub async fn skills_search_handler(
|
pub async fn skills_search_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(_user): AuthenticatedUser,
|
|
||||||
Json(req): Json<SkillSearchRequest>,
|
Json(req): Json<SkillSearchRequest>,
|
||||||
) -> Result<Json<SkillSearchResponse>, (StatusCode, String)> {
|
) -> Result<Json<SkillSearchResponse>, (StatusCode, String)> {
|
||||||
let registry = state.skill_registry.as_ref().ok_or((
|
let registry = state.skill_registry.as_ref().ok_or((
|
||||||
@@ -122,7 +119,6 @@ pub async fn skills_search_handler(
|
|||||||
|
|
||||||
pub async fn skills_install_handler(
|
pub async fn skills_install_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
headers: axum::http::HeaderMap,
|
headers: axum::http::HeaderMap,
|
||||||
Json(req): Json<SkillInstallRequest>,
|
Json(req): Json<SkillInstallRequest>,
|
||||||
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
|
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
|
||||||
@@ -139,8 +135,6 @@ pub async fn skills_install_handler(
|
|||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
tracing::info!(user_id = %user.user_id, skill = %req.name, "skill install requested");
|
|
||||||
|
|
||||||
let registry = state.skill_registry.as_ref().ok_or((
|
let registry = state.skill_registry.as_ref().ok_or((
|
||||||
StatusCode::NOT_IMPLEMENTED,
|
StatusCode::NOT_IMPLEMENTED,
|
||||||
"Skills system not enabled".to_string(),
|
"Skills system not enabled".to_string(),
|
||||||
@@ -225,7 +219,6 @@ pub async fn skills_install_handler(
|
|||||||
|
|
||||||
pub async fn skills_remove_handler(
|
pub async fn skills_remove_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(user): AuthenticatedUser,
|
|
||||||
headers: axum::http::HeaderMap,
|
headers: axum::http::HeaderMap,
|
||||||
Path(name): Path<String>,
|
Path(name): Path<String>,
|
||||||
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
|
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
|
||||||
@@ -241,8 +234,6 @@ pub async fn skills_remove_handler(
|
|||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
tracing::info!(user_id = %user.user_id, skill = %name, "skill remove requested");
|
|
||||||
|
|
||||||
let registry = state.skill_registry.as_ref().ok_or((
|
let registry = state.skill_registry.as_ref().ok_or((
|
||||||
StatusCode::NOT_IMPLEMENTED,
|
StatusCode::NOT_IMPLEMENTED,
|
||||||
"Skills system not enabled".to_string(),
|
"Skills system not enabled".to_string(),
|
||||||
|
|||||||
@@ -7,7 +7,6 @@ use axum::{
|
|||||||
};
|
};
|
||||||
|
|
||||||
use crate::bootstrap::ironclaw_base_dir;
|
use crate::bootstrap::ironclaw_base_dir;
|
||||||
use crate::channels::web::auth::AuthenticatedUser;
|
|
||||||
use crate::channels::web::types::*;
|
use crate::channels::web::types::*;
|
||||||
|
|
||||||
// --- Static file handlers ---
|
// --- Static file handlers ---
|
||||||
@@ -114,7 +113,6 @@ use crate::channels::web::server::GatewayState;
|
|||||||
|
|
||||||
pub async fn logs_events_handler(
|
pub async fn logs_events_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(_user): AuthenticatedUser,
|
|
||||||
) -> Result<
|
) -> Result<
|
||||||
Sse<impl futures::Stream<Item = Result<Event, Infallible>> + Send + 'static>,
|
Sse<impl futures::Stream<Item = Result<Event, Infallible>> + Send + 'static>,
|
||||||
(StatusCode, String),
|
(StatusCode, String),
|
||||||
@@ -154,7 +152,6 @@ pub async fn logs_events_handler(
|
|||||||
|
|
||||||
pub async fn gateway_status_handler(
|
pub async fn gateway_status_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
AuthenticatedUser(_user): AuthenticatedUser,
|
|
||||||
) -> Json<GatewayStatusResponse> {
|
) -> Json<GatewayStatusResponse> {
|
||||||
let sse_connections = state.sse.connection_count();
|
let sse_connections = state.sse.connection_count();
|
||||||
let ws_connections = state
|
let ws_connections = state
|
||||||
|
|||||||
+22
-90
@@ -31,9 +31,6 @@ pub mod ws;
|
|||||||
/// [`TestGatewayBuilder`](test_helpers::TestGatewayBuilder).
|
/// [`TestGatewayBuilder`](test_helpers::TestGatewayBuilder).
|
||||||
pub mod test_helpers;
|
pub mod test_helpers;
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests;
|
|
||||||
|
|
||||||
use std::net::SocketAddr;
|
use std::net::SocketAddr;
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
|
|
||||||
@@ -55,7 +52,6 @@ use crate::workspace::Workspace;
|
|||||||
|
|
||||||
use self::log_layer::{LogBroadcaster, LogLevelHandle};
|
use self::log_layer::{LogBroadcaster, LogLevelHandle};
|
||||||
|
|
||||||
use self::auth::MultiAuthState;
|
|
||||||
use self::server::GatewayState;
|
use self::server::GatewayState;
|
||||||
use self::sse::SseManager;
|
use self::sse::SseManager;
|
||||||
use self::types::SseEvent;
|
use self::types::SseEvent;
|
||||||
@@ -64,15 +60,14 @@ use self::types::SseEvent;
|
|||||||
pub struct GatewayChannel {
|
pub struct GatewayChannel {
|
||||||
config: GatewayConfig,
|
config: GatewayConfig,
|
||||||
state: Arc<GatewayState>,
|
state: Arc<GatewayState>,
|
||||||
/// Multi-user auth state (replaces bare auth_token).
|
/// The actual auth token in use (generated or from config).
|
||||||
auth: MultiAuthState,
|
auth_token: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl GatewayChannel {
|
impl GatewayChannel {
|
||||||
/// Create a new gateway channel.
|
/// Create a new gateway channel.
|
||||||
///
|
///
|
||||||
/// If no auth token is configured, generates a random one and prints it.
|
/// If no auth token is configured, generates a random one and prints it.
|
||||||
/// Builds a single-user `MultiAuthState` from the config.
|
|
||||||
pub fn new(config: GatewayConfig) -> Self {
|
pub fn new(config: GatewayConfig) -> Self {
|
||||||
let auth_token = config.auth_token.clone().unwrap_or_else(|| {
|
let auth_token = config.auth_token.clone().unwrap_or_else(|| {
|
||||||
use rand::RngCore;
|
use rand::RngCore;
|
||||||
@@ -82,13 +77,10 @@ impl GatewayChannel {
|
|||||||
bytes.iter().map(|b| format!("{b:02x}")).collect()
|
bytes.iter().map(|b| format!("{b:02x}")).collect()
|
||||||
});
|
});
|
||||||
|
|
||||||
let auth = MultiAuthState::single(auth_token, config.user_id.clone());
|
|
||||||
|
|
||||||
let state = Arc::new(GatewayState {
|
let state = Arc::new(GatewayState {
|
||||||
msg_tx: tokio::sync::RwLock::new(None),
|
msg_tx: tokio::sync::RwLock::new(None),
|
||||||
sse: Arc::new(SseManager::new()),
|
sse: SseManager::new(),
|
||||||
workspace: None,
|
workspace: None,
|
||||||
workspace_pool: None,
|
|
||||||
session_manager: None,
|
session_manager: None,
|
||||||
log_broadcaster: None,
|
log_broadcaster: None,
|
||||||
log_level_handle: None,
|
log_level_handle: None,
|
||||||
@@ -98,13 +90,13 @@ impl GatewayChannel {
|
|||||||
job_manager: None,
|
job_manager: None,
|
||||||
prompt_queue: None,
|
prompt_queue: None,
|
||||||
scheduler: None,
|
scheduler: None,
|
||||||
default_user_id: config.user_id.clone(),
|
user_id: config.user_id.clone(),
|
||||||
shutdown_tx: tokio::sync::RwLock::new(None),
|
shutdown_tx: tokio::sync::RwLock::new(None),
|
||||||
ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())),
|
ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())),
|
||||||
llm_provider: None,
|
llm_provider: None,
|
||||||
skill_registry: None,
|
skill_registry: None,
|
||||||
skill_catalog: None,
|
skill_catalog: None,
|
||||||
chat_rate_limiter: server::PerUserRateLimiter::new(30, 60),
|
chat_rate_limiter: server::RateLimiter::new(30, 60),
|
||||||
oauth_rate_limiter: server::RateLimiter::new(10, 60),
|
oauth_rate_limiter: server::RateLimiter::new(10, 60),
|
||||||
webhook_rate_limiter: server::RateLimiter::new(10, 60),
|
webhook_rate_limiter: server::RateLimiter::new(10, 60),
|
||||||
registry_entries: Vec::new(),
|
registry_entries: Vec::new(),
|
||||||
@@ -117,46 +109,7 @@ impl GatewayChannel {
|
|||||||
Self {
|
Self {
|
||||||
config,
|
config,
|
||||||
state,
|
state,
|
||||||
auth,
|
auth_token,
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Create a gateway channel with a pre-built multi-user auth state.
|
|
||||||
pub fn new_multi_auth(config: GatewayConfig, auth: MultiAuthState) -> Self {
|
|
||||||
let state = Arc::new(GatewayState {
|
|
||||||
msg_tx: tokio::sync::RwLock::new(None),
|
|
||||||
sse: Arc::new(SseManager::new()),
|
|
||||||
workspace: None,
|
|
||||||
workspace_pool: None,
|
|
||||||
session_manager: None,
|
|
||||||
log_broadcaster: None,
|
|
||||||
log_level_handle: None,
|
|
||||||
extension_manager: None,
|
|
||||||
tool_registry: None,
|
|
||||||
store: None,
|
|
||||||
job_manager: None,
|
|
||||||
prompt_queue: None,
|
|
||||||
scheduler: None,
|
|
||||||
default_user_id: config.user_id.clone(),
|
|
||||||
shutdown_tx: tokio::sync::RwLock::new(None),
|
|
||||||
ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())),
|
|
||||||
llm_provider: None,
|
|
||||||
skill_registry: None,
|
|
||||||
skill_catalog: None,
|
|
||||||
chat_rate_limiter: server::PerUserRateLimiter::new(30, 60),
|
|
||||||
oauth_rate_limiter: server::RateLimiter::new(10, 60),
|
|
||||||
registry_entries: Vec::new(),
|
|
||||||
cost_guard: None,
|
|
||||||
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
|
|
||||||
startup_time: std::time::Instant::now(),
|
|
||||||
webhook_rate_limiter: server::RateLimiter::new(10, 60),
|
|
||||||
active_config: server::ActiveConfigSnapshot::default(),
|
|
||||||
});
|
|
||||||
|
|
||||||
Self {
|
|
||||||
config,
|
|
||||||
state,
|
|
||||||
auth,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -165,9 +118,8 @@ impl GatewayChannel {
|
|||||||
let mut new_state = GatewayState {
|
let mut new_state = GatewayState {
|
||||||
msg_tx: tokio::sync::RwLock::new(None),
|
msg_tx: tokio::sync::RwLock::new(None),
|
||||||
// Preserve the existing broadcast channel so sender handles remain valid.
|
// Preserve the existing broadcast channel so sender handles remain valid.
|
||||||
sse: Arc::new(SseManager::from_sender(self.state.sse.sender())),
|
sse: SseManager::from_sender(self.state.sse.sender()),
|
||||||
workspace: self.state.workspace.clone(),
|
workspace: self.state.workspace.clone(),
|
||||||
workspace_pool: self.state.workspace_pool.clone(),
|
|
||||||
session_manager: self.state.session_manager.clone(),
|
session_manager: self.state.session_manager.clone(),
|
||||||
log_broadcaster: self.state.log_broadcaster.clone(),
|
log_broadcaster: self.state.log_broadcaster.clone(),
|
||||||
log_level_handle: self.state.log_level_handle.clone(),
|
log_level_handle: self.state.log_level_handle.clone(),
|
||||||
@@ -177,13 +129,13 @@ impl GatewayChannel {
|
|||||||
job_manager: self.state.job_manager.clone(),
|
job_manager: self.state.job_manager.clone(),
|
||||||
prompt_queue: self.state.prompt_queue.clone(),
|
prompt_queue: self.state.prompt_queue.clone(),
|
||||||
scheduler: self.state.scheduler.clone(),
|
scheduler: self.state.scheduler.clone(),
|
||||||
default_user_id: self.state.default_user_id.clone(),
|
user_id: self.state.user_id.clone(),
|
||||||
shutdown_tx: tokio::sync::RwLock::new(None),
|
shutdown_tx: tokio::sync::RwLock::new(None),
|
||||||
ws_tracker: self.state.ws_tracker.clone(),
|
ws_tracker: self.state.ws_tracker.clone(),
|
||||||
llm_provider: self.state.llm_provider.clone(),
|
llm_provider: self.state.llm_provider.clone(),
|
||||||
skill_registry: self.state.skill_registry.clone(),
|
skill_registry: self.state.skill_registry.clone(),
|
||||||
skill_catalog: self.state.skill_catalog.clone(),
|
skill_catalog: self.state.skill_catalog.clone(),
|
||||||
chat_rate_limiter: server::PerUserRateLimiter::new(30, 60),
|
chat_rate_limiter: server::RateLimiter::new(30, 60),
|
||||||
oauth_rate_limiter: server::RateLimiter::new(10, 60),
|
oauth_rate_limiter: server::RateLimiter::new(10, 60),
|
||||||
webhook_rate_limiter: server::RateLimiter::new(10, 60),
|
webhook_rate_limiter: server::RateLimiter::new(10, 60),
|
||||||
registry_entries: self.state.registry_entries.clone(),
|
registry_entries: self.state.registry_entries.clone(),
|
||||||
@@ -308,15 +260,9 @@ impl GatewayChannel {
|
|||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Inject the per-user workspace pool for multi-user mode.
|
/// Get the auth token (for printing to console on startup).
|
||||||
pub fn with_workspace_pool(mut self, pool: Arc<server::WorkspacePool>) -> Self {
|
|
||||||
self.rebuild_state(|s| s.workspace_pool = Some(pool));
|
|
||||||
self
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Get the first auth token (for printing to console on startup).
|
|
||||||
pub fn auth_token(&self) -> &str {
|
pub fn auth_token(&self) -> &str {
|
||||||
self.auth.first_token().unwrap_or("")
|
&self.auth_token
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Get a reference to the shared gateway state (for the agent to push SSE events).
|
/// Get a reference to the shared gateway state (for the agent to push SSE events).
|
||||||
@@ -345,7 +291,7 @@ impl Channel for GatewayChannel {
|
|||||||
),
|
),
|
||||||
})?;
|
})?;
|
||||||
|
|
||||||
server::start_server(addr, self.state.clone(), self.auth.clone()).await?;
|
server::start_server(addr, self.state.clone(), self.auth_token.clone()).await?;
|
||||||
|
|
||||||
Ok(Box::pin(ReceiverStream::new(rx)))
|
Ok(Box::pin(ReceiverStream::new(rx)))
|
||||||
}
|
}
|
||||||
@@ -365,13 +311,10 @@ impl Channel for GatewayChannel {
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
self.state.sse.broadcast_for_user(
|
self.state.sse.broadcast(SseEvent::Response {
|
||||||
&msg.user_id,
|
content: response.content,
|
||||||
SseEvent::Response {
|
thread_id,
|
||||||
content: response.content,
|
});
|
||||||
thread_id,
|
|
||||||
},
|
|
||||||
);
|
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
@@ -484,21 +427,13 @@ impl Channel for GatewayChannel {
|
|||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
// Scope events to the user when user_id is available in metadata.
|
self.state.sse.broadcast(event);
|
||||||
// When user_id is missing (heartbeat, routines), events go to all
|
|
||||||
// subscribers. In multi-tenant mode this leaks status across users.
|
|
||||||
if let Some(uid) = metadata.get("user_id").and_then(|v| v.as_str()) {
|
|
||||||
self.state.sse.broadcast_for_user(uid, event);
|
|
||||||
} else {
|
|
||||||
tracing::debug!("Status event missing user_id in metadata; broadcasting globally");
|
|
||||||
self.state.sse.broadcast(event);
|
|
||||||
}
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn broadcast(
|
async fn broadcast(
|
||||||
&self,
|
&self,
|
||||||
user_id: &str,
|
_user_id: &str,
|
||||||
response: OutgoingResponse,
|
response: OutgoingResponse,
|
||||||
) -> Result<(), ChannelError> {
|
) -> Result<(), ChannelError> {
|
||||||
let thread_id = match response.thread_id {
|
let thread_id = match response.thread_id {
|
||||||
@@ -510,13 +445,10 @@ impl Channel for GatewayChannel {
|
|||||||
return Ok(());
|
return Ok(());
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
self.state.sse.broadcast_for_user(
|
self.state.sse.broadcast(SseEvent::Response {
|
||||||
user_id,
|
content: response.content,
|
||||||
SseEvent::Response {
|
thread_id,
|
||||||
content: response.content,
|
});
|
||||||
thread_id,
|
|
||||||
},
|
|
||||||
);
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -463,10 +463,9 @@ fn build_tool_request(
|
|||||||
|
|
||||||
pub async fn chat_completions_handler(
|
pub async fn chat_completions_handler(
|
||||||
State(state): State<Arc<GatewayState>>,
|
State(state): State<Arc<GatewayState>>,
|
||||||
super::auth::AuthenticatedUser(user): super::auth::AuthenticatedUser,
|
|
||||||
Json(req): Json<OpenAiChatRequest>,
|
Json(req): Json<OpenAiChatRequest>,
|
||||||
) -> Result<impl IntoResponse, (StatusCode, Json<OpenAiErrorResponse>)> {
|
) -> Result<impl IntoResponse, (StatusCode, Json<OpenAiErrorResponse>)> {
|
||||||
if !state.chat_rate_limiter.check(&user.user_id) {
|
if !state.chat_rate_limiter.check() {
|
||||||
return Err(openai_error(
|
return Err(openai_error(
|
||||||
StatusCode::TOO_MANY_REQUESTS,
|
StatusCode::TOO_MANY_REQUESTS,
|
||||||
"Rate limit exceeded. Please try again later.",
|
"Rate limit exceeded. Please try again later.",
|
||||||
|
|||||||
+187
-509
File diff suppressed because it is too large
Load Diff
+28
-128
@@ -17,25 +17,9 @@ use crate::channels::web::types::SseEvent;
|
|||||||
/// Prevents resource exhaustion from connection flooding.
|
/// Prevents resource exhaustion from connection flooding.
|
||||||
const MAX_CONNECTIONS: u64 = 100;
|
const MAX_CONNECTIONS: u64 = 100;
|
||||||
|
|
||||||
/// Envelope for broadcast events: carries an optional user scope.
|
|
||||||
///
|
|
||||||
/// `user_id = None` means the event is global (e.g. Heartbeat) and delivered
|
|
||||||
/// to all subscribers. `user_id = Some(id)` means the event is only delivered
|
|
||||||
/// to subscribers that match that user_id.
|
|
||||||
#[derive(Debug, Clone)]
|
|
||||||
pub(crate) struct ScopedEvent {
|
|
||||||
pub(crate) user_id: Option<String>,
|
|
||||||
pub(crate) event: SseEvent,
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Manages SSE broadcast to all connected browser tabs.
|
/// Manages SSE broadcast to all connected browser tabs.
|
||||||
///
|
|
||||||
/// In multi-user mode, events are scoped by user_id so that each subscriber
|
|
||||||
/// only receives events intended for their user (plus global events like
|
|
||||||
/// Heartbeat). In single-user mode, all events are delivered to all subscribers
|
|
||||||
/// (backwards compatible).
|
|
||||||
pub struct SseManager {
|
pub struct SseManager {
|
||||||
tx: broadcast::Sender<ScopedEvent>,
|
tx: broadcast::Sender<SseEvent>,
|
||||||
connection_count: Arc<AtomicU64>,
|
connection_count: Arc<AtomicU64>,
|
||||||
max_connections: u64,
|
max_connections: u64,
|
||||||
}
|
}
|
||||||
@@ -61,7 +45,7 @@ impl SseManager {
|
|||||||
/// only be called before the server starts accepting connections (i.e.,
|
/// only be called before the server starts accepting connections (i.e.,
|
||||||
/// during startup wiring). Calling it after connections are established
|
/// during startup wiring). Calling it after connections are established
|
||||||
/// will break connection tracking and allow exceeding `MAX_CONNECTIONS`.
|
/// will break connection tracking and allow exceeding `MAX_CONNECTIONS`.
|
||||||
pub(crate) fn from_sender(tx: broadcast::Sender<ScopedEvent>) -> Self {
|
pub fn from_sender(tx: broadcast::Sender<SseEvent>) -> Self {
|
||||||
Self {
|
Self {
|
||||||
tx,
|
tx,
|
||||||
connection_count: Arc::new(AtomicU64::new(0)),
|
connection_count: Arc::new(AtomicU64::new(0)),
|
||||||
@@ -69,28 +53,15 @@ impl SseManager {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Get a clone of the broadcast sender for use by other components.
|
/// Broadcast an event to all connected clients.
|
||||||
pub(crate) fn sender(&self) -> broadcast::Sender<ScopedEvent> {
|
|
||||||
self.tx.clone()
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Broadcast an event to all connected clients (global/unscoped).
|
|
||||||
pub fn broadcast(&self, event: SseEvent) {
|
pub fn broadcast(&self, event: SseEvent) {
|
||||||
let _ = self.tx.send(ScopedEvent {
|
// Ignore send errors (no receivers is fine)
|
||||||
user_id: None,
|
let _ = self.tx.send(event);
|
||||||
event,
|
|
||||||
});
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Broadcast an event scoped to a specific user.
|
/// Get a clone of the broadcast sender for use by other components.
|
||||||
///
|
pub fn sender(&self) -> broadcast::Sender<SseEvent> {
|
||||||
/// Only subscribers for this user_id (or unscoped subscribers) will
|
self.tx.clone()
|
||||||
/// receive the event.
|
|
||||||
pub fn broadcast_for_user(&self, user_id: &str, event: SseEvent) {
|
|
||||||
let _ = self.tx.send(ScopedEvent {
|
|
||||||
user_id: Some(user_id.to_string()),
|
|
||||||
event,
|
|
||||||
});
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Get current number of active connections.
|
/// Get current number of active connections.
|
||||||
@@ -100,15 +71,11 @@ impl SseManager {
|
|||||||
|
|
||||||
/// Create a raw broadcast subscription for non-SSE consumers (e.g. WebSocket).
|
/// Create a raw broadcast subscription for non-SSE consumers (e.g. WebSocket).
|
||||||
///
|
///
|
||||||
/// When `user_id` is `Some`, only events scoped to that user (or global
|
/// Returns a stream of `SseEvent` values and increments/decrements the
|
||||||
/// events) are delivered. When `None`, all events are delivered (single-user
|
/// connection counter on creation/drop, just like `subscribe()` does for SSE.
|
||||||
/// backwards compatibility).
|
|
||||||
///
|
///
|
||||||
/// Returns `None` if the maximum connection limit has been reached.
|
/// Returns `None` if the maximum connection limit has been reached.
|
||||||
pub fn subscribe_raw(
|
pub fn subscribe_raw(&self) -> Option<impl Stream<Item = SseEvent> + Send + 'static + use<>> {
|
||||||
&self,
|
|
||||||
user_id: Option<String>,
|
|
||||||
) -> Option<impl Stream<Item = SseEvent> + Send + 'static + use<>> {
|
|
||||||
// Atomically increment only if below the limit. This prevents
|
// Atomically increment only if below the limit. This prevents
|
||||||
// concurrent callers from overshooting max_connections.
|
// concurrent callers from overshooting max_connections.
|
||||||
let counter = Arc::clone(&self.connection_count);
|
let counter = Arc::clone(&self.connection_count);
|
||||||
@@ -124,19 +91,7 @@ impl SseManager {
|
|||||||
.ok()?;
|
.ok()?;
|
||||||
let rx = self.tx.subscribe();
|
let rx = self.tx.subscribe();
|
||||||
|
|
||||||
let stream = BroadcastStream::new(rx).filter_map(move |result| match result {
|
let stream = BroadcastStream::new(rx).filter_map(|result| result.ok());
|
||||||
Ok(scoped) => {
|
|
||||||
// Global events (user_id=None) always pass through.
|
|
||||||
// Scoped events only pass if the subscriber matches (or subscriber is unscoped).
|
|
||||||
match (&user_id, &scoped.user_id) {
|
|
||||||
(_, None) => Some(scoped.event), // global -> all
|
|
||||||
(None, _) => Some(scoped.event), // unscoped subscriber -> all
|
|
||||||
(Some(sub), Some(ev)) if sub == ev => Some(scoped.event), // match
|
|
||||||
_ => None, // different user -> skip
|
|
||||||
}
|
|
||||||
}
|
|
||||||
Err(_) => None,
|
|
||||||
});
|
|
||||||
|
|
||||||
Some(CountedStream {
|
Some(CountedStream {
|
||||||
inner: stream,
|
inner: stream,
|
||||||
@@ -146,13 +101,9 @@ impl SseManager {
|
|||||||
|
|
||||||
/// Create a new SSE stream for a client connection.
|
/// Create a new SSE stream for a client connection.
|
||||||
///
|
///
|
||||||
/// When `user_id` is `Some`, only events for that user (or global events)
|
|
||||||
/// are delivered. When `None`, all events are delivered.
|
|
||||||
///
|
|
||||||
/// Returns `None` if the maximum connection limit has been reached.
|
/// Returns `None` if the maximum connection limit has been reached.
|
||||||
pub fn subscribe(
|
pub fn subscribe(
|
||||||
&self,
|
&self,
|
||||||
user_id: Option<String>,
|
|
||||||
) -> Option<Sse<impl Stream<Item = Result<Event, Infallible>> + Send + 'static + use<>>> {
|
) -> Option<Sse<impl Stream<Item = Result<Event, Infallible>> + Send + 'static + use<>>> {
|
||||||
// Atomically increment only if below the limit.
|
// Atomically increment only if below the limit.
|
||||||
let counter = Arc::clone(&self.connection_count);
|
let counter = Arc::clone(&self.connection_count);
|
||||||
@@ -169,23 +120,9 @@ impl SseManager {
|
|||||||
let rx = self.tx.subscribe();
|
let rx = self.tx.subscribe();
|
||||||
|
|
||||||
let stream = BroadcastStream::new(rx)
|
let stream = BroadcastStream::new(rx)
|
||||||
.filter_map(move |result| match result {
|
.filter_map(|result| result.ok())
|
||||||
Ok(scoped) => match (&user_id, &scoped.user_id) {
|
.map(|event| {
|
||||||
(_, None) => Some(scoped.event),
|
let data = serde_json::to_string(&event).unwrap_or_default();
|
||||||
(None, _) => Some(scoped.event),
|
|
||||||
(Some(sub), Some(ev)) if sub == ev => Some(scoped.event),
|
|
||||||
_ => None,
|
|
||||||
},
|
|
||||||
Err(_) => None,
|
|
||||||
})
|
|
||||||
.filter_map(|event| {
|
|
||||||
let data = match serde_json::to_string(&event) {
|
|
||||||
Ok(s) => s,
|
|
||||||
Err(e) => {
|
|
||||||
tracing::warn!("Failed to serialize SSE event: {}", e);
|
|
||||||
return None;
|
|
||||||
}
|
|
||||||
};
|
|
||||||
let event_type = match &event {
|
let event_type = match &event {
|
||||||
SseEvent::Response { .. } => "response",
|
SseEvent::Response { .. } => "response",
|
||||||
SseEvent::Thinking { .. } => "thinking",
|
SseEvent::Thinking { .. } => "thinking",
|
||||||
@@ -210,7 +147,7 @@ impl SseManager {
|
|||||||
SseEvent::TurnCost { .. } => "turn_cost",
|
SseEvent::TurnCost { .. } => "turn_cost",
|
||||||
SseEvent::ExtensionStatus { .. } => "extension_status",
|
SseEvent::ExtensionStatus { .. } => "extension_status",
|
||||||
};
|
};
|
||||||
Some(Ok(Event::default().event(event_type).data(data)))
|
Ok(Event::default().event(event_type).data(data))
|
||||||
});
|
});
|
||||||
|
|
||||||
// Wrap in a stream that decrements on drop
|
// Wrap in a stream that decrements on drop
|
||||||
@@ -278,14 +215,16 @@ mod tests {
|
|||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_broadcast_to_receiver() {
|
async fn test_broadcast_to_receiver() {
|
||||||
let manager = SseManager::new();
|
let manager = SseManager::new();
|
||||||
let mut stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe"));
|
let mut rx = BroadcastStream::new(manager.tx.subscribe());
|
||||||
|
|
||||||
manager.broadcast(SseEvent::Status {
|
manager.broadcast(SseEvent::Status {
|
||||||
message: "test".to_string(),
|
message: "test".to_string(),
|
||||||
thread_id: None,
|
thread_id: None,
|
||||||
});
|
});
|
||||||
|
|
||||||
let event = stream.next().await.unwrap();
|
let event = rx.next().await;
|
||||||
|
assert!(event.is_some());
|
||||||
|
let event = event.unwrap().unwrap();
|
||||||
match event {
|
match event {
|
||||||
SseEvent::Status { message, .. } => assert_eq!(message, "test"),
|
SseEvent::Status { message, .. } => assert_eq!(message, "test"),
|
||||||
_ => panic!("unexpected event type"),
|
_ => panic!("unexpected event type"),
|
||||||
@@ -295,7 +234,7 @@ mod tests {
|
|||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_subscribe_raw_receives_events() {
|
async fn test_subscribe_raw_receives_events() {
|
||||||
let manager = SseManager::new();
|
let manager = SseManager::new();
|
||||||
let mut stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe"));
|
let mut stream = Box::pin(manager.subscribe_raw().expect("should subscribe"));
|
||||||
|
|
||||||
assert_eq!(manager.connection_count(), 1);
|
assert_eq!(manager.connection_count(), 1);
|
||||||
|
|
||||||
@@ -315,7 +254,7 @@ mod tests {
|
|||||||
async fn test_subscribe_raw_decrements_on_drop() {
|
async fn test_subscribe_raw_decrements_on_drop() {
|
||||||
let manager = SseManager::new();
|
let manager = SseManager::new();
|
||||||
{
|
{
|
||||||
let _stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe"));
|
let _stream = Box::pin(manager.subscribe_raw().expect("should subscribe"));
|
||||||
assert_eq!(manager.connection_count(), 1);
|
assert_eq!(manager.connection_count(), 1);
|
||||||
}
|
}
|
||||||
// Stream dropped, counter should decrement
|
// Stream dropped, counter should decrement
|
||||||
@@ -325,8 +264,8 @@ mod tests {
|
|||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_subscribe_raw_multiple_subscribers() {
|
async fn test_subscribe_raw_multiple_subscribers() {
|
||||||
let manager = SseManager::new();
|
let manager = SseManager::new();
|
||||||
let mut s1 = Box::pin(manager.subscribe_raw(None).expect("should subscribe"));
|
let mut s1 = Box::pin(manager.subscribe_raw().expect("should subscribe"));
|
||||||
let mut s2 = Box::pin(manager.subscribe_raw(None).expect("should subscribe"));
|
let mut s2 = Box::pin(manager.subscribe_raw().expect("should subscribe"));
|
||||||
assert_eq!(manager.connection_count(), 2);
|
assert_eq!(manager.connection_count(), 2);
|
||||||
|
|
||||||
manager.broadcast(SseEvent::Heartbeat);
|
manager.broadcast(SseEvent::Heartbeat);
|
||||||
@@ -347,51 +286,12 @@ mod tests {
|
|||||||
let mut manager = SseManager::new();
|
let mut manager = SseManager::new();
|
||||||
manager.max_connections = 2; // Low limit for testing
|
manager.max_connections = 2; // Low limit for testing
|
||||||
|
|
||||||
let _s1 = Box::pin(manager.subscribe_raw(None).expect("first should succeed"));
|
let _s1 = Box::pin(manager.subscribe_raw().expect("first should succeed"));
|
||||||
let _s2 = Box::pin(manager.subscribe_raw(None).expect("second should succeed"));
|
let _s2 = Box::pin(manager.subscribe_raw().expect("second should succeed"));
|
||||||
assert_eq!(manager.connection_count(), 2);
|
assert_eq!(manager.connection_count(), 2);
|
||||||
|
|
||||||
// Third should be rejected
|
// Third should be rejected
|
||||||
assert!(manager.subscribe_raw(None).is_none());
|
assert!(manager.subscribe_raw().is_none());
|
||||||
assert!(manager.subscribe(None).is_none());
|
assert!(manager.subscribe().is_none());
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_scoped_events_filtered_by_user() {
|
|
||||||
let manager = SseManager::new();
|
|
||||||
let mut alice = Box::pin(
|
|
||||||
manager
|
|
||||||
.subscribe_raw(Some("alice".to_string()))
|
|
||||||
.expect("subscribe"),
|
|
||||||
);
|
|
||||||
let mut bob = Box::pin(
|
|
||||||
manager
|
|
||||||
.subscribe_raw(Some("bob".to_string()))
|
|
||||||
.expect("subscribe"),
|
|
||||||
);
|
|
||||||
|
|
||||||
// Send event scoped to alice
|
|
||||||
manager.broadcast_for_user(
|
|
||||||
"alice",
|
|
||||||
SseEvent::Status {
|
|
||||||
message: "alice only".to_string(),
|
|
||||||
thread_id: None,
|
|
||||||
},
|
|
||||||
);
|
|
||||||
|
|
||||||
// Send global event
|
|
||||||
manager.broadcast(SseEvent::Heartbeat);
|
|
||||||
|
|
||||||
// Alice gets her scoped event
|
|
||||||
let e = alice.next().await.unwrap();
|
|
||||||
assert!(matches!(e, SseEvent::Status { .. }));
|
|
||||||
|
|
||||||
// Alice also gets the global heartbeat
|
|
||||||
let e = alice.next().await.unwrap();
|
|
||||||
assert!(matches!(e, SseEvent::Heartbeat));
|
|
||||||
|
|
||||||
// Bob only gets the global heartbeat (alice's event was filtered)
|
|
||||||
let e = bob.next().await.unwrap();
|
|
||||||
assert!(matches!(e, SseEvent::Heartbeat));
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,8 +10,7 @@ use std::sync::Arc;
|
|||||||
use tokio::sync::mpsc;
|
use tokio::sync::mpsc;
|
||||||
|
|
||||||
use crate::channels::IncomingMessage;
|
use crate::channels::IncomingMessage;
|
||||||
use crate::channels::web::auth::MultiAuthState;
|
use crate::channels::web::server::{GatewayState, RateLimiter, start_server};
|
||||||
use crate::channels::web::server::{GatewayState, PerUserRateLimiter, RateLimiter, start_server};
|
|
||||||
use crate::channels::web::sse::SseManager;
|
use crate::channels::web::sse::SseManager;
|
||||||
use crate::channels::web::ws::WsConnectionTracker;
|
use crate::channels::web::ws::WsConnectionTracker;
|
||||||
|
|
||||||
@@ -65,9 +64,8 @@ impl TestGatewayBuilder {
|
|||||||
pub fn build(self) -> Arc<GatewayState> {
|
pub fn build(self) -> Arc<GatewayState> {
|
||||||
Arc::new(GatewayState {
|
Arc::new(GatewayState {
|
||||||
msg_tx: tokio::sync::RwLock::new(self.msg_tx),
|
msg_tx: tokio::sync::RwLock::new(self.msg_tx),
|
||||||
sse: Arc::new(SseManager::new()),
|
sse: SseManager::new(),
|
||||||
workspace: None,
|
workspace: None,
|
||||||
workspace_pool: None,
|
|
||||||
session_manager: None,
|
session_manager: None,
|
||||||
log_broadcaster: None,
|
log_broadcaster: None,
|
||||||
log_level_handle: None,
|
log_level_handle: None,
|
||||||
@@ -76,14 +74,14 @@ impl TestGatewayBuilder {
|
|||||||
store: None,
|
store: None,
|
||||||
job_manager: None,
|
job_manager: None,
|
||||||
prompt_queue: None,
|
prompt_queue: None,
|
||||||
default_user_id: self.user_id,
|
user_id: self.user_id,
|
||||||
shutdown_tx: tokio::sync::RwLock::new(None),
|
shutdown_tx: tokio::sync::RwLock::new(None),
|
||||||
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
|
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
|
||||||
llm_provider: self.llm_provider,
|
llm_provider: self.llm_provider,
|
||||||
skill_registry: None,
|
skill_registry: None,
|
||||||
skill_catalog: None,
|
skill_catalog: None,
|
||||||
scheduler: None,
|
scheduler: None,
|
||||||
chat_rate_limiter: PerUserRateLimiter::new(30, 60),
|
chat_rate_limiter: RateLimiter::new(30, 60),
|
||||||
oauth_rate_limiter: RateLimiter::new(10, 60),
|
oauth_rate_limiter: RateLimiter::new(10, 60),
|
||||||
webhook_rate_limiter: RateLimiter::new(10, 60),
|
webhook_rate_limiter: RateLimiter::new(10, 60),
|
||||||
registry_entries: Vec::new(),
|
registry_entries: Vec::new(),
|
||||||
@@ -100,26 +98,11 @@ impl TestGatewayBuilder {
|
|||||||
self,
|
self,
|
||||||
auth_token: &str,
|
auth_token: &str,
|
||||||
) -> Result<(SocketAddr, Arc<GatewayState>), crate::error::ChannelError> {
|
) -> Result<(SocketAddr, Arc<GatewayState>), crate::error::ChannelError> {
|
||||||
let auth = MultiAuthState::single(auth_token.to_string(), "test-user".to_string());
|
|
||||||
let state = self.build();
|
let state = self.build();
|
||||||
let addr: SocketAddr = "127.0.0.1:0"
|
let addr: SocketAddr = "127.0.0.1:0"
|
||||||
.parse()
|
.parse()
|
||||||
.expect("hard-coded address must parse"); // safety: constant literal
|
.expect("hard-coded address must parse");
|
||||||
let bound = start_server(addr, state.clone(), auth).await?;
|
let bound = start_server(addr, state.clone(), auth_token.to_string()).await?;
|
||||||
Ok((bound, state))
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Build the state and start a gateway server with multi-user auth.
|
|
||||||
/// Returns the bound address and the shared state.
|
|
||||||
pub async fn start_multi(
|
|
||||||
self,
|
|
||||||
auth: MultiAuthState,
|
|
||||||
) -> Result<(SocketAddr, Arc<GatewayState>), crate::error::ChannelError> {
|
|
||||||
let state = self.build();
|
|
||||||
let addr: SocketAddr = "127.0.0.1:0"
|
|
||||||
.parse()
|
|
||||||
.expect("hard-coded address must parse"); // safety: constant literal
|
|
||||||
let bound = start_server(addr, state.clone(), auth).await?;
|
|
||||||
Ok((bound, state))
|
Ok((bound, state))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,3 +0,0 @@
|
|||||||
//! Integration tests for the web gateway module.
|
|
||||||
|
|
||||||
mod multi_tenant;
|
|
||||||
@@ -1,796 +0,0 @@
|
|||||||
//! Multi-tenant isolation tests for the web gateway.
|
|
||||||
//!
|
|
||||||
//! Tests cover workspace pool scoping, job handler isolation, and auth
|
|
||||||
//! enforcement on protected endpoints. Uses `LibSqlBackend::new_local()`
|
|
||||||
//! with a temporary directory for a real (but ephemeral) database.
|
|
||||||
|
|
||||||
use std::collections::HashMap;
|
|
||||||
use std::sync::Arc;
|
|
||||||
use std::time::Duration;
|
|
||||||
|
|
||||||
use axum::Router;
|
|
||||||
use axum::body::Body;
|
|
||||||
use axum::http::{Method, Request, StatusCode};
|
|
||||||
use axum::middleware;
|
|
||||||
use axum::routing::{delete, get, post};
|
|
||||||
use tower::ServiceExt;
|
|
||||||
use uuid::Uuid;
|
|
||||||
|
|
||||||
use crate::channels::web::auth::{
|
|
||||||
AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware,
|
|
||||||
};
|
|
||||||
use crate::channels::web::server::{
|
|
||||||
ActiveConfigSnapshot, GatewayState, PerUserRateLimiter, PromptQueue, RateLimiter, WorkspacePool,
|
|
||||||
};
|
|
||||||
use crate::channels::web::sse::SseManager;
|
|
||||||
|
|
||||||
// ── Helpers ────────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
/// Create a two-user `MultiAuthState` for alice and bob.
|
|
||||||
fn two_user_auth() -> MultiAuthState {
|
|
||||||
let mut tokens = HashMap::new();
|
|
||||||
tokens.insert(
|
|
||||||
"tok-alice".to_string(),
|
|
||||||
UserIdentity {
|
|
||||||
user_id: "alice".to_string(),
|
|
||||||
workspace_read_scopes: vec!["shared".to_string()],
|
|
||||||
},
|
|
||||||
);
|
|
||||||
tokens.insert(
|
|
||||||
"tok-bob".to_string(),
|
|
||||||
UserIdentity {
|
|
||||||
user_id: "bob".to_string(),
|
|
||||||
workspace_read_scopes: vec!["shared".to_string(), "alice".to_string()],
|
|
||||||
},
|
|
||||||
);
|
|
||||||
MultiAuthState::multi(tokens)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Build a `GatewayState` with configurable store and prompt queue.
|
|
||||||
fn build_state(
|
|
||||||
store: Option<Arc<dyn crate::db::Database>>,
|
|
||||||
prompt_queue: Option<PromptQueue>,
|
|
||||||
) -> Arc<GatewayState> {
|
|
||||||
Arc::new(GatewayState {
|
|
||||||
msg_tx: tokio::sync::RwLock::new(None),
|
|
||||||
sse: Arc::new(SseManager::new()),
|
|
||||||
workspace: None,
|
|
||||||
workspace_pool: None,
|
|
||||||
session_manager: None,
|
|
||||||
log_broadcaster: None,
|
|
||||||
log_level_handle: None,
|
|
||||||
extension_manager: None,
|
|
||||||
tool_registry: None,
|
|
||||||
store,
|
|
||||||
job_manager: None,
|
|
||||||
prompt_queue,
|
|
||||||
default_user_id: "test".to_string(),
|
|
||||||
shutdown_tx: tokio::sync::RwLock::new(None),
|
|
||||||
ws_tracker: None,
|
|
||||||
llm_provider: None,
|
|
||||||
skill_registry: None,
|
|
||||||
skill_catalog: None,
|
|
||||||
scheduler: None,
|
|
||||||
chat_rate_limiter: PerUserRateLimiter::new(30, 60),
|
|
||||||
oauth_rate_limiter: RateLimiter::new(10, 60),
|
|
||||||
webhook_rate_limiter: RateLimiter::new(10, 60),
|
|
||||||
registry_entries: Vec::new(),
|
|
||||||
cost_guard: None,
|
|
||||||
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
|
|
||||||
startup_time: std::time::Instant::now(),
|
|
||||||
active_config: ActiveConfigSnapshot::default(),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Create a libSQL-backed test database in a temporary directory.
|
|
||||||
///
|
|
||||||
/// Returns the database and a `TempDir` guard — the database file is
|
|
||||||
/// deleted when the guard is dropped.
|
|
||||||
#[cfg(feature = "libsql")]
|
|
||||||
async fn test_db() -> (Arc<dyn crate::db::Database>, tempfile::TempDir) {
|
|
||||||
use crate::db::Database;
|
|
||||||
let dir = tempfile::tempdir().expect("failed to create temp dir"); // safety: test-only
|
|
||||||
let path = dir.path().join("test.db");
|
|
||||||
let backend = crate::db::libsql::LibSqlBackend::new_local(&path)
|
|
||||||
.await
|
|
||||||
.expect("failed to create test LibSqlBackend"); // safety: test-only
|
|
||||||
backend
|
|
||||||
.run_migrations()
|
|
||||||
.await
|
|
||||||
.expect("failed to run migrations"); // safety: test-only
|
|
||||||
(Arc::new(backend) as Arc<dyn crate::db::Database>, dir)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Build a minimal Routine for testing.
|
|
||||||
fn make_routine(user_id: &str, name: &str) -> crate::agent::routine::Routine {
|
|
||||||
let now = chrono::Utc::now();
|
|
||||||
crate::agent::routine::Routine {
|
|
||||||
id: Uuid::new_v4(),
|
|
||||||
name: name.to_string(),
|
|
||||||
description: format!("Test routine: {name}"),
|
|
||||||
user_id: user_id.to_string(),
|
|
||||||
enabled: true,
|
|
||||||
trigger: crate::agent::routine::Trigger::Cron {
|
|
||||||
schedule: "0 9 * * *".to_string(),
|
|
||||||
timezone: None,
|
|
||||||
},
|
|
||||||
action: crate::agent::routine::RoutineAction::Lightweight {
|
|
||||||
prompt: "hello".to_string(),
|
|
||||||
context_paths: vec![],
|
|
||||||
max_tokens: 1024,
|
|
||||||
use_tools: false,
|
|
||||||
max_tool_rounds: 3,
|
|
||||||
},
|
|
||||||
guardrails: crate::agent::routine::RoutineGuardrails {
|
|
||||||
cooldown: Duration::from_secs(60),
|
|
||||||
max_concurrent: 1,
|
|
||||||
dedup_window: None,
|
|
||||||
},
|
|
||||||
notify: crate::agent::routine::NotifyConfig {
|
|
||||||
channel: None,
|
|
||||||
user: None,
|
|
||||||
on_success: false,
|
|
||||||
on_failure: true,
|
|
||||||
on_attention: true,
|
|
||||||
},
|
|
||||||
last_run_at: None,
|
|
||||||
next_fire_at: None,
|
|
||||||
run_count: 0,
|
|
||||||
consecutive_failures: 0,
|
|
||||||
state: serde_json::json!({}),
|
|
||||||
created_at: now,
|
|
||||||
updated_at: now,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Build a minimal SandboxJobRecord for testing.
|
|
||||||
fn make_sandbox_job(user_id: &str, task: &str) -> crate::history::SandboxJobRecord {
|
|
||||||
let now = chrono::Utc::now();
|
|
||||||
crate::history::SandboxJobRecord {
|
|
||||||
id: Uuid::new_v4(),
|
|
||||||
task: task.to_string(),
|
|
||||||
status: "completed".to_string(),
|
|
||||||
user_id: user_id.to_string(),
|
|
||||||
project_dir: format!("/tmp/test-{}", Uuid::new_v4()),
|
|
||||||
success: Some(true),
|
|
||||||
failure_reason: None,
|
|
||||||
created_at: now,
|
|
||||||
started_at: Some(now),
|
|
||||||
completed_at: Some(now),
|
|
||||||
credential_grants_json: "[]".to_string(),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ═══════════════════════════════════════════════════════════════════════
|
|
||||||
// WorkspacePool Tests
|
|
||||||
// ═══════════════════════════════════════════════════════════════════════
|
|
||||||
|
|
||||||
#[cfg(feature = "libsql")]
|
|
||||||
mod workspace_pool {
|
|
||||||
use super::*;
|
|
||||||
use crate::config::{WorkspaceConfig, WorkspaceSearchConfig};
|
|
||||||
use crate::workspace::EmbeddingCacheConfig;
|
|
||||||
use crate::workspace::layer::MemoryLayer;
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_workspace_pool_applies_search_config() {
|
|
||||||
let (db, _dir) = test_db().await;
|
|
||||||
let search_config = WorkspaceSearchConfig {
|
|
||||||
rrf_k: 42,
|
|
||||||
..Default::default()
|
|
||||||
};
|
|
||||||
let pool = WorkspacePool::new(
|
|
||||||
db,
|
|
||||||
None,
|
|
||||||
EmbeddingCacheConfig::default(),
|
|
||||||
search_config,
|
|
||||||
WorkspaceConfig::default(),
|
|
||||||
);
|
|
||||||
let identity = UserIdentity {
|
|
||||||
user_id: "alice".to_string(),
|
|
||||||
workspace_read_scopes: vec![],
|
|
||||||
};
|
|
||||||
let ws = pool.get_or_create(&identity).await;
|
|
||||||
assert_eq!(ws.user_id(), "alice");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_workspace_pool_applies_memory_layers() {
|
|
||||||
let (db, _dir) = test_db().await;
|
|
||||||
let layers = vec![MemoryLayer {
|
|
||||||
name: "shared-layer".to_string(),
|
|
||||||
scope: "shared".to_string(),
|
|
||||||
writable: false,
|
|
||||||
sensitivity: Default::default(),
|
|
||||||
}];
|
|
||||||
let ws_config = WorkspaceConfig {
|
|
||||||
memory_layers: layers,
|
|
||||||
read_scopes: vec![],
|
|
||||||
};
|
|
||||||
let pool = WorkspacePool::new(
|
|
||||||
db,
|
|
||||||
None,
|
|
||||||
EmbeddingCacheConfig::default(),
|
|
||||||
WorkspaceSearchConfig::default(),
|
|
||||||
ws_config,
|
|
||||||
);
|
|
||||||
let identity = UserIdentity {
|
|
||||||
user_id: "alice".to_string(),
|
|
||||||
workspace_read_scopes: vec![],
|
|
||||||
};
|
|
||||||
let ws = pool.get_or_create(&identity).await;
|
|
||||||
// Memory layer scope "shared" should appear in read_user_ids.
|
|
||||||
assert!(
|
|
||||||
ws.read_user_ids().contains(&"shared".to_string()),
|
|
||||||
"expected 'shared' in read_user_ids, got {:?}",
|
|
||||||
ws.read_user_ids()
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_workspace_pool_applies_identity_read_scopes() {
|
|
||||||
let (db, _dir) = test_db().await;
|
|
||||||
let pool = WorkspacePool::new(
|
|
||||||
db,
|
|
||||||
None,
|
|
||||||
EmbeddingCacheConfig::default(),
|
|
||||||
WorkspaceSearchConfig::default(),
|
|
||||||
WorkspaceConfig::default(),
|
|
||||||
);
|
|
||||||
let identity = UserIdentity {
|
|
||||||
user_id: "bob".to_string(),
|
|
||||||
workspace_read_scopes: vec!["alice".to_string(), "shared".to_string()],
|
|
||||||
};
|
|
||||||
let ws = pool.get_or_create(&identity).await;
|
|
||||||
assert_eq!(ws.user_id(), "bob");
|
|
||||||
assert!(
|
|
||||||
ws.read_user_ids().contains(&"alice".to_string()),
|
|
||||||
"expected 'alice' in read_user_ids from identity scopes"
|
|
||||||
);
|
|
||||||
assert!(
|
|
||||||
ws.read_user_ids().contains(&"shared".to_string()),
|
|
||||||
"expected 'shared' in read_user_ids from identity scopes"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_workspace_pool_caches_per_user() {
|
|
||||||
let (db, _dir) = test_db().await;
|
|
||||||
let pool = WorkspacePool::new(
|
|
||||||
db,
|
|
||||||
None,
|
|
||||||
EmbeddingCacheConfig::default(),
|
|
||||||
WorkspaceSearchConfig::default(),
|
|
||||||
WorkspaceConfig::default(),
|
|
||||||
);
|
|
||||||
let alice_id = UserIdentity {
|
|
||||||
user_id: "alice".to_string(),
|
|
||||||
workspace_read_scopes: vec![],
|
|
||||||
};
|
|
||||||
let bob_id = UserIdentity {
|
|
||||||
user_id: "bob".to_string(),
|
|
||||||
workspace_read_scopes: vec![],
|
|
||||||
};
|
|
||||||
|
|
||||||
let alice_ws1 = pool.get_or_create(&alice_id).await;
|
|
||||||
let alice_ws2 = pool.get_or_create(&alice_id).await;
|
|
||||||
let bob_ws = pool.get_or_create(&bob_id).await;
|
|
||||||
|
|
||||||
// Same user gets the same Arc.
|
|
||||||
assert!(Arc::ptr_eq(&alice_ws1, &alice_ws2));
|
|
||||||
// Different users get different instances.
|
|
||||||
assert!(!Arc::ptr_eq(&alice_ws1, &bob_ws));
|
|
||||||
assert_eq!(alice_ws1.user_id(), "alice");
|
|
||||||
assert_eq!(bob_ws.user_id(), "bob");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_workspace_pool_combines_global_and_identity_scopes() {
|
|
||||||
let (db, _dir) = test_db().await;
|
|
||||||
let ws_config = WorkspaceConfig {
|
|
||||||
memory_layers: vec![],
|
|
||||||
read_scopes: vec!["global-shared".to_string()],
|
|
||||||
};
|
|
||||||
let pool = WorkspacePool::new(
|
|
||||||
db,
|
|
||||||
None,
|
|
||||||
EmbeddingCacheConfig::default(),
|
|
||||||
WorkspaceSearchConfig::default(),
|
|
||||||
ws_config,
|
|
||||||
);
|
|
||||||
let identity = UserIdentity {
|
|
||||||
user_id: "alice".to_string(),
|
|
||||||
workspace_read_scopes: vec!["token-scope".to_string()],
|
|
||||||
};
|
|
||||||
let ws = pool.get_or_create(&identity).await;
|
|
||||||
let scopes = ws.read_user_ids();
|
|
||||||
// Primary scope
|
|
||||||
assert!(scopes.contains(&"alice".to_string()));
|
|
||||||
// Global config scope
|
|
||||||
assert!(
|
|
||||||
scopes.contains(&"global-shared".to_string()),
|
|
||||||
"expected global scope 'global-shared', got {:?}",
|
|
||||||
scopes
|
|
||||||
);
|
|
||||||
// Token identity scope
|
|
||||||
assert!(
|
|
||||||
scopes.contains(&"token-scope".to_string()),
|
|
||||||
"expected token scope 'token-scope', got {:?}",
|
|
||||||
scopes
|
|
||||||
);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ═══════════════════════════════════════════════════════════════════════
|
|
||||||
// Jobs Handler Isolation Tests
|
|
||||||
// ═══════════════════════════════════════════════════════════════════════
|
|
||||||
|
|
||||||
#[cfg(feature = "libsql")]
|
|
||||||
mod jobs_isolation {
|
|
||||||
use super::*;
|
|
||||||
use crate::channels::web::handlers::jobs::{
|
|
||||||
jobs_cancel_handler, jobs_prompt_handler, jobs_restart_handler, jobs_summary_handler,
|
|
||||||
};
|
|
||||||
// SandboxStore methods are accessed through the Database supertrait.
|
|
||||||
|
|
||||||
/// Build a router with job endpoints behind multi-user auth.
|
|
||||||
fn jobs_router(state: Arc<GatewayState>, auth: MultiAuthState) -> Router {
|
|
||||||
Router::new()
|
|
||||||
.route("/api/jobs/summary", get(jobs_summary_handler))
|
|
||||||
.route("/api/jobs/{id}/cancel", post(jobs_cancel_handler))
|
|
||||||
.route("/api/jobs/{id}/restart", post(jobs_restart_handler))
|
|
||||||
.route("/api/jobs/{id}/prompt", post(jobs_prompt_handler))
|
|
||||||
.layer(middleware::from_fn_with_state(auth, auth_middleware))
|
|
||||||
.with_state(state)
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_jobs_summary_scoped_to_user() {
|
|
||||||
let (db, _dir) = test_db().await;
|
|
||||||
|
|
||||||
// Insert sandbox jobs for alice and bob.
|
|
||||||
let alice_job = make_sandbox_job("alice", "alice task");
|
|
||||||
let bob_job = make_sandbox_job("bob", "bob task");
|
|
||||||
db.save_sandbox_job(&alice_job).await.unwrap();
|
|
||||||
db.save_sandbox_job(&bob_job).await.unwrap();
|
|
||||||
|
|
||||||
let state = build_state(Some(db), None);
|
|
||||||
let auth = two_user_auth();
|
|
||||||
let app = jobs_router(state, auth);
|
|
||||||
|
|
||||||
// Alice should see 1 job.
|
|
||||||
let req = Request::builder()
|
|
||||||
.uri("/api/jobs/summary")
|
|
||||||
.header("Authorization", "Bearer tok-alice")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.clone().oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(resp.status(), StatusCode::OK);
|
|
||||||
let body: serde_json::Value =
|
|
||||||
serde_json::from_slice(&axum::body::to_bytes(resp.into_body(), 4096).await.unwrap())
|
|
||||||
.unwrap();
|
|
||||||
assert_eq!(body["total"], 1, "alice should see only her own jobs");
|
|
||||||
|
|
||||||
// Bob should see 1 job.
|
|
||||||
let req = Request::builder()
|
|
||||||
.uri("/api/jobs/summary")
|
|
||||||
.header("Authorization", "Bearer tok-bob")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(resp.status(), StatusCode::OK);
|
|
||||||
let body: serde_json::Value =
|
|
||||||
serde_json::from_slice(&axum::body::to_bytes(resp.into_body(), 4096).await.unwrap())
|
|
||||||
.unwrap();
|
|
||||||
assert_eq!(body["total"], 1, "bob should see only his own jobs");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_jobs_restart_rejects_other_user() {
|
|
||||||
let (db, _dir) = test_db().await;
|
|
||||||
|
|
||||||
// Insert a failed sandbox job owned by alice.
|
|
||||||
let mut alice_job = make_sandbox_job("alice", "alice task");
|
|
||||||
alice_job.status = "failed".to_string();
|
|
||||||
alice_job.success = Some(false);
|
|
||||||
db.save_sandbox_job(&alice_job).await.unwrap();
|
|
||||||
|
|
||||||
let state = build_state(Some(db), None);
|
|
||||||
let auth = two_user_auth();
|
|
||||||
let app = jobs_router(state, auth);
|
|
||||||
|
|
||||||
// Bob tries to restart alice's job.
|
|
||||||
let req = Request::builder()
|
|
||||||
.method(Method::POST)
|
|
||||||
.uri(format!("/api/jobs/{}/restart", alice_job.id))
|
|
||||||
.header("Authorization", "Bearer tok-bob")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(
|
|
||||||
resp.status(),
|
|
||||||
StatusCode::NOT_FOUND,
|
|
||||||
"bob should not be able to restart alice's job"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_jobs_prompt_works_for_agent_jobs() {
|
|
||||||
let (db, _dir) = test_db().await;
|
|
||||||
|
|
||||||
// Insert a running sandbox job owned by alice in claude_code mode.
|
|
||||||
let mut alice_job = make_sandbox_job("alice", "prompt test");
|
|
||||||
alice_job.status = "running".to_string();
|
|
||||||
alice_job.success = None;
|
|
||||||
alice_job.completed_at = None;
|
|
||||||
db.save_sandbox_job(&alice_job).await.unwrap();
|
|
||||||
db.update_sandbox_job_mode(alice_job.id, "claude_code")
|
|
||||||
.await
|
|
||||||
.unwrap();
|
|
||||||
|
|
||||||
let prompt_queue: PromptQueue =
|
|
||||||
Arc::new(tokio::sync::Mutex::new(std::collections::HashMap::new()));
|
|
||||||
let state = build_state(Some(db), Some(prompt_queue.clone()));
|
|
||||||
let auth = two_user_auth();
|
|
||||||
let app = jobs_router(state, auth);
|
|
||||||
|
|
||||||
// Alice prompts her own job.
|
|
||||||
let req = Request::builder()
|
|
||||||
.method(Method::POST)
|
|
||||||
.uri(format!("/api/jobs/{}/prompt", alice_job.id))
|
|
||||||
.header("Authorization", "Bearer tok-alice")
|
|
||||||
.header("Content-Type", "application/json")
|
|
||||||
.body(Body::from(
|
|
||||||
serde_json::to_string(&serde_json::json!({"content": "hello"})).unwrap(),
|
|
||||||
))
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(
|
|
||||||
resp.status(),
|
|
||||||
StatusCode::OK,
|
|
||||||
"alice should be able to prompt her own job"
|
|
||||||
);
|
|
||||||
|
|
||||||
// Verify prompt was enqueued.
|
|
||||||
let queue = prompt_queue.lock().await;
|
|
||||||
assert!(
|
|
||||||
queue.contains_key(&alice_job.id),
|
|
||||||
"prompt queue should contain alice's job"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_jobs_prompt_rejects_other_user() {
|
|
||||||
let (db, _dir) = test_db().await;
|
|
||||||
|
|
||||||
let mut alice_job = make_sandbox_job("alice", "alice task");
|
|
||||||
alice_job.status = "running".to_string();
|
|
||||||
alice_job.success = None;
|
|
||||||
alice_job.completed_at = None;
|
|
||||||
db.save_sandbox_job(&alice_job).await.unwrap();
|
|
||||||
db.update_sandbox_job_mode(alice_job.id, "claude_code")
|
|
||||||
.await
|
|
||||||
.unwrap();
|
|
||||||
|
|
||||||
let prompt_queue: PromptQueue =
|
|
||||||
Arc::new(tokio::sync::Mutex::new(std::collections::HashMap::new()));
|
|
||||||
let state = build_state(Some(db), Some(prompt_queue));
|
|
||||||
let auth = two_user_auth();
|
|
||||||
let app = jobs_router(state, auth);
|
|
||||||
|
|
||||||
// Bob tries to prompt alice's job.
|
|
||||||
let req = Request::builder()
|
|
||||||
.method(Method::POST)
|
|
||||||
.uri(format!("/api/jobs/{}/prompt", alice_job.id))
|
|
||||||
.header("Authorization", "Bearer tok-bob")
|
|
||||||
.header("Content-Type", "application/json")
|
|
||||||
.body(Body::from(
|
|
||||||
serde_json::to_string(&serde_json::json!({"content": "sneaky"})).unwrap(),
|
|
||||||
))
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(
|
|
||||||
resp.status(),
|
|
||||||
StatusCode::NOT_FOUND,
|
|
||||||
"bob should not be able to prompt alice's job"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_jobs_cancel_rejects_other_user() {
|
|
||||||
let (db, _dir) = test_db().await;
|
|
||||||
|
|
||||||
let mut alice_job = make_sandbox_job("alice", "alice running");
|
|
||||||
alice_job.status = "running".to_string();
|
|
||||||
alice_job.success = None;
|
|
||||||
alice_job.completed_at = None;
|
|
||||||
db.save_sandbox_job(&alice_job).await.unwrap();
|
|
||||||
|
|
||||||
let state = build_state(Some(db), None);
|
|
||||||
let auth = two_user_auth();
|
|
||||||
let app = jobs_router(state, auth);
|
|
||||||
|
|
||||||
// Bob tries to cancel alice's job.
|
|
||||||
let req = Request::builder()
|
|
||||||
.method(Method::POST)
|
|
||||||
.uri(format!("/api/jobs/{}/cancel", alice_job.id))
|
|
||||||
.header("Authorization", "Bearer tok-bob")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(
|
|
||||||
resp.status(),
|
|
||||||
StatusCode::NOT_FOUND,
|
|
||||||
"bob should not be able to cancel alice's job"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ═══════════════════════════════════════════════════════════════════════
|
|
||||||
// Routines Isolation Tests
|
|
||||||
// ═══════════════════════════════════════════════════════════════════════
|
|
||||||
|
|
||||||
#[cfg(feature = "libsql")]
|
|
||||||
mod routines_isolation {
|
|
||||||
use super::*;
|
|
||||||
use crate::channels::web::handlers::routines::{
|
|
||||||
routines_delete_handler, routines_detail_handler, routines_list_handler,
|
|
||||||
routines_summary_handler, routines_toggle_handler,
|
|
||||||
};
|
|
||||||
// RoutineStore methods are accessed through the Database supertrait.
|
|
||||||
|
|
||||||
fn routines_router(state: Arc<GatewayState>, auth: MultiAuthState) -> Router {
|
|
||||||
Router::new()
|
|
||||||
.route("/api/routines", get(routines_list_handler))
|
|
||||||
.route("/api/routines/summary", get(routines_summary_handler))
|
|
||||||
.route("/api/routines/{id}", get(routines_detail_handler))
|
|
||||||
.route("/api/routines/{id}/toggle", post(routines_toggle_handler))
|
|
||||||
.route("/api/routines/{id}", delete(routines_delete_handler))
|
|
||||||
.layer(middleware::from_fn_with_state(auth, auth_middleware))
|
|
||||||
.with_state(state)
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_routines_isolation() {
|
|
||||||
let (db, _dir) = test_db().await;
|
|
||||||
|
|
||||||
// Create routines for alice and bob.
|
|
||||||
let alice_routine = make_routine("alice", "alice-daily");
|
|
||||||
let bob_routine = make_routine("bob", "bob-daily");
|
|
||||||
db.create_routine(&alice_routine).await.unwrap();
|
|
||||||
db.create_routine(&bob_routine).await.unwrap();
|
|
||||||
|
|
||||||
let state = build_state(Some(db), None);
|
|
||||||
let auth = two_user_auth();
|
|
||||||
let app = routines_router(state, auth);
|
|
||||||
|
|
||||||
// Alice sees only her routine in the list.
|
|
||||||
let req = Request::builder()
|
|
||||||
.uri("/api/routines")
|
|
||||||
.header("Authorization", "Bearer tok-alice")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.clone().oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(resp.status(), StatusCode::OK);
|
|
||||||
let body: serde_json::Value =
|
|
||||||
serde_json::from_slice(&axum::body::to_bytes(resp.into_body(), 8192).await.unwrap())
|
|
||||||
.unwrap();
|
|
||||||
let routines = body["routines"].as_array().unwrap();
|
|
||||||
assert_eq!(routines.len(), 1, "alice should see only her routines");
|
|
||||||
assert_eq!(routines[0]["name"], "alice-daily");
|
|
||||||
|
|
||||||
// Bob sees only his routine.
|
|
||||||
let req = Request::builder()
|
|
||||||
.uri("/api/routines")
|
|
||||||
.header("Authorization", "Bearer tok-bob")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.clone().oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(resp.status(), StatusCode::OK);
|
|
||||||
let body: serde_json::Value =
|
|
||||||
serde_json::from_slice(&axum::body::to_bytes(resp.into_body(), 8192).await.unwrap())
|
|
||||||
.unwrap();
|
|
||||||
let routines = body["routines"].as_array().unwrap();
|
|
||||||
assert_eq!(routines.len(), 1, "bob should see only his routines");
|
|
||||||
assert_eq!(routines[0]["name"], "bob-daily");
|
|
||||||
|
|
||||||
// Bob cannot view alice's routine detail.
|
|
||||||
let req = Request::builder()
|
|
||||||
.uri(format!("/api/routines/{}", alice_routine.id))
|
|
||||||
.header("Authorization", "Bearer tok-bob")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.clone().oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(
|
|
||||||
resp.status(),
|
|
||||||
StatusCode::NOT_FOUND,
|
|
||||||
"bob should not see alice's routine detail"
|
|
||||||
);
|
|
||||||
|
|
||||||
// Bob cannot toggle alice's routine.
|
|
||||||
let req = Request::builder()
|
|
||||||
.method(Method::POST)
|
|
||||||
.uri(format!("/api/routines/{}/toggle", alice_routine.id))
|
|
||||||
.header("Authorization", "Bearer tok-bob")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.clone().oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(
|
|
||||||
resp.status(),
|
|
||||||
StatusCode::NOT_FOUND,
|
|
||||||
"bob should not toggle alice's routine"
|
|
||||||
);
|
|
||||||
|
|
||||||
// Bob cannot delete alice's routine.
|
|
||||||
let req = Request::builder()
|
|
||||||
.method(Method::DELETE)
|
|
||||||
.uri(format!("/api/routines/{}", alice_routine.id))
|
|
||||||
.header("Authorization", "Bearer tok-bob")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(
|
|
||||||
resp.status(),
|
|
||||||
StatusCode::NOT_FOUND,
|
|
||||||
"bob should not delete alice's routine"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ═══════════════════════════════════════════════════════════════════════
|
|
||||||
// Handler Auth Enforcement Tests
|
|
||||||
// ═══════════════════════════════════════════════════════════════════════
|
|
||||||
|
|
||||||
mod auth_enforcement {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
/// Dummy handler that extracts `AuthenticatedUser` — if the auth middleware
|
|
||||||
/// rejects the request, this handler is never reached.
|
|
||||||
async fn authed_handler(AuthenticatedUser(_user): AuthenticatedUser) -> &'static str {
|
|
||||||
"ok"
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Build a router with the real auth middleware and dummy handlers at all
|
|
||||||
/// the paths we want to verify require authentication.
|
|
||||||
fn auth_test_router(auth: MultiAuthState) -> Router {
|
|
||||||
let state = build_state(None, None);
|
|
||||||
Router::new()
|
|
||||||
// Routines
|
|
||||||
.route("/api/routines", get(authed_handler))
|
|
||||||
.route("/api/routines/summary", get(authed_handler))
|
|
||||||
.route("/api/routines/{id}", get(authed_handler))
|
|
||||||
.route("/api/routines/{id}/toggle", post(authed_handler))
|
|
||||||
.route("/api/routines/{id}", delete(authed_handler))
|
|
||||||
// Skills
|
|
||||||
.route("/api/skills", get(authed_handler))
|
|
||||||
.route("/api/skills/search", post(authed_handler))
|
|
||||||
.route("/api/skills/install", post(authed_handler))
|
|
||||||
.route("/api/skills/{name}", delete(authed_handler))
|
|
||||||
// Logs
|
|
||||||
.route("/api/logs/events", get(authed_handler))
|
|
||||||
.route("/api/logs/level", get(authed_handler).put(authed_handler))
|
|
||||||
// Gateway status
|
|
||||||
.route("/api/gateway/status", get(authed_handler))
|
|
||||||
.layer(middleware::from_fn_with_state(auth, auth_middleware))
|
|
||||||
.with_state(state)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Send a request without auth and assert it returns UNAUTHORIZED.
|
|
||||||
async fn assert_requires_auth(app: &Router, method: Method, uri: &str) {
|
|
||||||
let req = Request::builder()
|
|
||||||
.method(method.clone())
|
|
||||||
.uri(uri)
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.clone().oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(
|
|
||||||
resp.status(),
|
|
||||||
StatusCode::UNAUTHORIZED,
|
|
||||||
"{} {} should require auth",
|
|
||||||
method,
|
|
||||||
uri
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Send a request with a valid token and assert it succeeds.
|
|
||||||
async fn assert_passes_with_token(app: &Router, method: Method, uri: &str, token: &str) {
|
|
||||||
let req = Request::builder()
|
|
||||||
.method(method.clone())
|
|
||||||
.uri(uri)
|
|
||||||
.header("Authorization", format!("Bearer {token}"))
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.clone().oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(
|
|
||||||
resp.status(),
|
|
||||||
StatusCode::OK,
|
|
||||||
"{} {} should pass with valid token",
|
|
||||||
method,
|
|
||||||
uri
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_routines_handlers_require_auth() {
|
|
||||||
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
|
|
||||||
let app = auth_test_router(auth);
|
|
||||||
let id = Uuid::new_v4();
|
|
||||||
|
|
||||||
assert_requires_auth(&app, Method::GET, "/api/routines").await;
|
|
||||||
assert_requires_auth(&app, Method::GET, "/api/routines/summary").await;
|
|
||||||
assert_requires_auth(&app, Method::GET, &format!("/api/routines/{id}")).await;
|
|
||||||
assert_requires_auth(&app, Method::POST, &format!("/api/routines/{id}/toggle")).await;
|
|
||||||
assert_requires_auth(&app, Method::DELETE, &format!("/api/routines/{id}")).await;
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_skills_handlers_require_auth() {
|
|
||||||
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
|
|
||||||
let app = auth_test_router(auth);
|
|
||||||
|
|
||||||
assert_requires_auth(&app, Method::GET, "/api/skills").await;
|
|
||||||
assert_requires_auth(&app, Method::POST, "/api/skills/search").await;
|
|
||||||
assert_requires_auth(&app, Method::POST, "/api/skills/install").await;
|
|
||||||
assert_requires_auth(&app, Method::DELETE, "/api/skills/test-skill").await;
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_logs_handlers_require_auth() {
|
|
||||||
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
|
|
||||||
let app = auth_test_router(auth);
|
|
||||||
|
|
||||||
assert_requires_auth(&app, Method::GET, "/api/logs/events").await;
|
|
||||||
assert_requires_auth(&app, Method::GET, "/api/logs/level").await;
|
|
||||||
assert_requires_auth(&app, Method::PUT, "/api/logs/level").await;
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_gateway_status_requires_auth() {
|
|
||||||
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
|
|
||||||
let app = auth_test_router(auth);
|
|
||||||
|
|
||||||
assert_requires_auth(&app, Method::GET, "/api/gateway/status").await;
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_valid_token_passes_all_endpoints() {
|
|
||||||
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
|
|
||||||
let app = auth_test_router(auth);
|
|
||||||
let id = Uuid::new_v4();
|
|
||||||
|
|
||||||
assert_passes_with_token(&app, Method::GET, "/api/routines", "secret-tok").await;
|
|
||||||
assert_passes_with_token(&app, Method::GET, "/api/skills", "secret-tok").await;
|
|
||||||
assert_passes_with_token(&app, Method::GET, "/api/logs/events", "secret-tok").await;
|
|
||||||
assert_passes_with_token(&app, Method::GET, "/api/gateway/status", "secret-tok").await;
|
|
||||||
assert_passes_with_token(
|
|
||||||
&app,
|
|
||||||
Method::GET,
|
|
||||||
&format!("/api/routines/{id}"),
|
|
||||||
"secret-tok",
|
|
||||||
)
|
|
||||||
.await;
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_wrong_token_rejected_on_all_endpoints() {
|
|
||||||
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
|
|
||||||
let app = auth_test_router(auth);
|
|
||||||
|
|
||||||
// Wrong token should be rejected.
|
|
||||||
let req = Request::builder()
|
|
||||||
.uri("/api/routines")
|
|
||||||
.header("Authorization", "Bearer wrong-tok")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.clone().oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
|
|
||||||
|
|
||||||
let req = Request::builder()
|
|
||||||
.uri("/api/gateway/status")
|
|
||||||
.header("Authorization", "Bearer wrong-tok")
|
|
||||||
.body(Body::empty())
|
|
||||||
.unwrap();
|
|
||||||
let resp = app.oneshot(req).await.unwrap();
|
|
||||||
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+13
-24
@@ -62,11 +62,7 @@ impl Default for WsConnectionTracker {
|
|||||||
///
|
///
|
||||||
/// When either task ends (client disconnect or broadcast closed), both are
|
/// When either task ends (client disconnect or broadcast closed), both are
|
||||||
/// cleaned up.
|
/// cleaned up.
|
||||||
pub async fn handle_ws_connection(
|
pub async fn handle_ws_connection(socket: WebSocket, state: Arc<GatewayState>) {
|
||||||
socket: WebSocket,
|
|
||||||
state: Arc<GatewayState>,
|
|
||||||
user: crate::channels::web::auth::UserIdentity,
|
|
||||||
) {
|
|
||||||
let (mut ws_sink, mut ws_stream) = socket.split();
|
let (mut ws_sink, mut ws_stream) = socket.split();
|
||||||
|
|
||||||
// Track connection
|
// Track connection
|
||||||
@@ -75,9 +71,9 @@ pub async fn handle_ws_connection(
|
|||||||
}
|
}
|
||||||
let tracker_for_drop = state.ws_tracker.clone();
|
let tracker_for_drop = state.ws_tracker.clone();
|
||||||
|
|
||||||
// Subscribe to broadcast events (same source as SSE), scoped to this user.
|
// Subscribe to broadcast events (same source as SSE).
|
||||||
// Reject if we've hit the connection limit.
|
// Reject if we've hit the connection limit.
|
||||||
let Some(raw_stream) = state.sse.subscribe_raw(Some(user.user_id.clone())) else {
|
let Some(raw_stream) = state.sse.subscribe_raw() else {
|
||||||
tracing::warn!("WebSocket rejected: too many connections");
|
tracing::warn!("WebSocket rejected: too many connections");
|
||||||
// Decrement the WS tracker we already incremented above.
|
// Decrement the WS tracker we already incremented above.
|
||||||
if let Some(ref tracker) = tracker_for_drop {
|
if let Some(ref tracker) = tracker_for_drop {
|
||||||
@@ -121,7 +117,7 @@ pub async fn handle_ws_connection(
|
|||||||
});
|
});
|
||||||
|
|
||||||
// Receiver task: read client frames and route to agent
|
// Receiver task: read client frames and route to agent
|
||||||
let user_id = user.user_id;
|
let user_id = state.user_id.clone();
|
||||||
while let Some(Ok(frame)) = ws_stream.next().await {
|
while let Some(Ok(frame)) = ws_stream.next().await {
|
||||||
match frame {
|
match frame {
|
||||||
Message::Text(text) => {
|
Message::Text(text) => {
|
||||||
@@ -267,14 +263,10 @@ async fn handle_client_message(
|
|||||||
token,
|
token,
|
||||||
} => {
|
} => {
|
||||||
if let Some(ref ext_mgr) = state.extension_manager {
|
if let Some(ref ext_mgr) = state.extension_manager {
|
||||||
match ext_mgr
|
match ext_mgr.configure_token(&extension_name, &token).await {
|
||||||
.configure_token(&extension_name, &token, user_id)
|
|
||||||
.await
|
|
||||||
{
|
|
||||||
Ok(result) => {
|
Ok(result) => {
|
||||||
if result.verification.is_some() {
|
if result.verification.is_some() {
|
||||||
state.sse.broadcast_for_user(
|
state.sse.broadcast(
|
||||||
user_id,
|
|
||||||
crate::channels::web::types::SseEvent::AuthRequired {
|
crate::channels::web::types::SseEvent::AuthRequired {
|
||||||
extension_name: extension_name.clone(),
|
extension_name: extension_name.clone(),
|
||||||
instructions: Some(result.message),
|
instructions: Some(result.message),
|
||||||
@@ -283,9 +275,8 @@ async fn handle_client_message(
|
|||||||
},
|
},
|
||||||
);
|
);
|
||||||
} else {
|
} else {
|
||||||
crate::channels::web::server::clear_auth_mode(state, user_id).await;
|
crate::channels::web::server::clear_auth_mode(state).await;
|
||||||
state.sse.broadcast_for_user(
|
state.sse.broadcast(
|
||||||
user_id,
|
|
||||||
crate::channels::web::types::SseEvent::AuthCompleted {
|
crate::channels::web::types::SseEvent::AuthCompleted {
|
||||||
extension_name,
|
extension_name,
|
||||||
success: true,
|
success: true,
|
||||||
@@ -297,8 +288,7 @@ async fn handle_client_message(
|
|||||||
Err(e) => {
|
Err(e) => {
|
||||||
let msg = format!("Auth failed: {}", e);
|
let msg = format!("Auth failed: {}", e);
|
||||||
if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) {
|
if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) {
|
||||||
state.sse.broadcast_for_user(
|
state.sse.broadcast(
|
||||||
user_id,
|
|
||||||
crate::channels::web::types::SseEvent::AuthRequired {
|
crate::channels::web::types::SseEvent::AuthRequired {
|
||||||
extension_name: extension_name.clone(),
|
extension_name: extension_name.clone(),
|
||||||
instructions: Some(msg.clone()),
|
instructions: Some(msg.clone()),
|
||||||
@@ -321,7 +311,7 @@ async fn handle_client_message(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
WsClientMessage::AuthCancel { .. } => {
|
WsClientMessage::AuthCancel { .. } => {
|
||||||
crate::channels::web::server::clear_auth_mode(state, user_id).await;
|
crate::channels::web::server::clear_auth_mode(state).await;
|
||||||
}
|
}
|
||||||
WsClientMessage::Ping => {
|
WsClientMessage::Ping => {
|
||||||
let _ = direct_tx.send(WsServerMessage::Pong).await;
|
let _ = direct_tx.send(WsServerMessage::Pong).await;
|
||||||
@@ -508,9 +498,8 @@ mod tests {
|
|||||||
|
|
||||||
GatewayState {
|
GatewayState {
|
||||||
msg_tx: tokio::sync::RwLock::new(msg_tx),
|
msg_tx: tokio::sync::RwLock::new(msg_tx),
|
||||||
sse: Arc::new(SseManager::new()),
|
sse: SseManager::new(),
|
||||||
workspace: None,
|
workspace: None,
|
||||||
workspace_pool: None,
|
|
||||||
session_manager: None,
|
session_manager: None,
|
||||||
log_broadcaster: None,
|
log_broadcaster: None,
|
||||||
log_level_handle: None,
|
log_level_handle: None,
|
||||||
@@ -520,13 +509,13 @@ mod tests {
|
|||||||
job_manager: None,
|
job_manager: None,
|
||||||
prompt_queue: None,
|
prompt_queue: None,
|
||||||
scheduler: None,
|
scheduler: None,
|
||||||
default_user_id: "test".to_string(),
|
user_id: "test".to_string(),
|
||||||
shutdown_tx: tokio::sync::RwLock::new(None),
|
shutdown_tx: tokio::sync::RwLock::new(None),
|
||||||
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
|
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
|
||||||
llm_provider: None,
|
llm_provider: None,
|
||||||
skill_registry: None,
|
skill_registry: None,
|
||||||
skill_catalog: None,
|
skill_catalog: None,
|
||||||
chat_rate_limiter: crate::channels::web::server::PerUserRateLimiter::new(30, 60),
|
chat_rate_limiter: crate::channels::web::server::RateLimiter::new(30, 60),
|
||||||
oauth_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60),
|
oauth_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60),
|
||||||
webhook_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60),
|
webhook_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60),
|
||||||
registry_entries: Vec::new(),
|
registry_entries: Vec::new(),
|
||||||
|
|||||||
@@ -25,7 +25,6 @@ pub mod import;
|
|||||||
mod logs;
|
mod logs;
|
||||||
mod mcp;
|
mod mcp;
|
||||||
pub mod memory;
|
pub mod memory;
|
||||||
mod models;
|
|
||||||
pub mod oauth_defaults;
|
pub mod oauth_defaults;
|
||||||
mod pairing;
|
mod pairing;
|
||||||
mod registry;
|
mod registry;
|
||||||
@@ -46,7 +45,6 @@ pub use logs::{LogsCommand, run_logs_command};
|
|||||||
pub use mcp::{McpCommand, run_mcp_command};
|
pub use mcp::{McpCommand, run_mcp_command};
|
||||||
pub use memory::MemoryCommand;
|
pub use memory::MemoryCommand;
|
||||||
pub use memory::run_memory_command_with_db;
|
pub use memory::run_memory_command_with_db;
|
||||||
pub use models::{ModelsCommand, run_models_command};
|
|
||||||
pub use pairing::{PairingCommand, run_pairing_command, run_pairing_command_with_store};
|
pub use pairing::{PairingCommand, run_pairing_command, run_pairing_command_with_store};
|
||||||
pub use registry::{RegistryCommand, run_registry_command};
|
pub use registry::{RegistryCommand, run_registry_command};
|
||||||
pub use routines::{RoutinesCommand, run_routines_command};
|
pub use routines::{RoutinesCommand, run_routines_command};
|
||||||
@@ -219,14 +217,6 @@ pub enum Command {
|
|||||||
)]
|
)]
|
||||||
Hooks(HooksCommand),
|
Hooks(HooksCommand),
|
||||||
|
|
||||||
/// Manage LLM providers and models
|
|
||||||
#[command(
|
|
||||||
subcommand,
|
|
||||||
about = "Manage LLM providers and models",
|
|
||||||
long_about = "List providers, view current configuration, and set active provider/model.\nExamples:\n ironclaw models list\n ironclaw models list openai --verbose\n ironclaw models status\n ironclaw models set gpt-4o\n ironclaw models set-provider anthropic --model claude-sonnet-4-6-20250514"
|
|
||||||
)]
|
|
||||||
Models(ModelsCommand),
|
|
||||||
|
|
||||||
/// Probe external dependencies and validate configuration
|
/// Probe external dependencies and validate configuration
|
||||||
#[command(
|
#[command(
|
||||||
about = "Run diagnostics",
|
about = "Run diagnostics",
|
||||||
|
|||||||
@@ -1,864 +0,0 @@
|
|||||||
//! Models management CLI commands.
|
|
||||||
//!
|
|
||||||
//! Provides subcommands for listing providers, viewing current model
|
|
||||||
//! configuration, and setting the active provider/model. Settings are
|
|
||||||
//! persisted to both `config.toml` and `~/.ironclaw/.env` so changes
|
|
||||||
//! take effect immediately (no DB connection required).
|
|
||||||
|
|
||||||
use clap::Subcommand;
|
|
||||||
use std::path::Path;
|
|
||||||
|
|
||||||
use crate::llm::registry::ProviderRegistry;
|
|
||||||
use crate::settings::Settings;
|
|
||||||
|
|
||||||
#[derive(Subcommand, Debug, Clone)]
|
|
||||||
pub enum ModelsCommand {
|
|
||||||
/// List providers (or available models for a specific provider)
|
|
||||||
List {
|
|
||||||
/// Show only a specific provider (by ID or alias)
|
|
||||||
provider: Option<String>,
|
|
||||||
|
|
||||||
/// Show detailed information (env vars, base URL, protocol)
|
|
||||||
#[arg(short, long)]
|
|
||||||
verbose: bool,
|
|
||||||
|
|
||||||
/// Output as JSON
|
|
||||||
#[arg(long)]
|
|
||||||
json: bool,
|
|
||||||
},
|
|
||||||
|
|
||||||
/// Show current model configuration
|
|
||||||
Status {
|
|
||||||
/// Output as JSON
|
|
||||||
#[arg(long)]
|
|
||||||
json: bool,
|
|
||||||
},
|
|
||||||
|
|
||||||
/// Set the default model
|
|
||||||
Set {
|
|
||||||
/// Model name (e.g., "gpt-5-mini", "claude-sonnet-4-6-20250514")
|
|
||||||
model: String,
|
|
||||||
},
|
|
||||||
|
|
||||||
/// Set the LLM provider
|
|
||||||
SetProvider {
|
|
||||||
/// Provider ID or alias (e.g., "openai", "anthropic", "ollama")
|
|
||||||
provider: String,
|
|
||||||
|
|
||||||
/// Also set the model (defaults to provider's default model)
|
|
||||||
#[arg(long)]
|
|
||||||
model: Option<String>,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Run the models CLI subcommand.
|
|
||||||
pub async fn run_models_command(
|
|
||||||
cmd: ModelsCommand,
|
|
||||||
config_path: Option<&Path>,
|
|
||||||
) -> anyhow::Result<()> {
|
|
||||||
match cmd {
|
|
||||||
ModelsCommand::List {
|
|
||||||
provider,
|
|
||||||
verbose,
|
|
||||||
json,
|
|
||||||
} => {
|
|
||||||
if let Some(ref id) = provider {
|
|
||||||
cmd_show_provider(id, verbose, json, config_path).await
|
|
||||||
} else {
|
|
||||||
cmd_list_providers(verbose, json, config_path).await
|
|
||||||
}
|
|
||||||
}
|
|
||||||
ModelsCommand::Status { json } => cmd_status(json, config_path),
|
|
||||||
ModelsCommand::Set { model } => cmd_set_model(&model, config_path),
|
|
||||||
ModelsCommand::SetProvider { provider, model } => {
|
|
||||||
cmd_set_provider(&provider, model.as_deref(), config_path)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ─── Shared helpers ───────────────────────────────────────────────
|
|
||||||
|
|
||||||
/// Resolve the currently active backend and model from env + settings.
|
|
||||||
fn resolve_active(config_path: Option<&Path>) -> (String, String) {
|
|
||||||
let settings = load_settings(config_path);
|
|
||||||
resolve_active_from_settings(&settings)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Resolve active backend + model from a pre-loaded Settings.
|
|
||||||
fn resolve_active_from_settings(settings: &Settings) -> (String, String) {
|
|
||||||
let backend = std::env::var("LLM_BACKEND")
|
|
||||||
.ok()
|
|
||||||
.or_else(|| settings.llm_backend.clone())
|
|
||||||
.unwrap_or_else(|| "nearai".to_string());
|
|
||||||
|
|
||||||
let registry = ProviderRegistry::load();
|
|
||||||
|
|
||||||
let canonical_backend = registry
|
|
||||||
.find(&backend)
|
|
||||||
.map(|d| d.id.clone())
|
|
||||||
.unwrap_or_else(|| backend.clone());
|
|
||||||
|
|
||||||
let model = if canonical_backend == "nearai" {
|
|
||||||
std::env::var("NEARAI_MODEL")
|
|
||||||
.ok()
|
|
||||||
.or_else(|| settings.selected_model.clone())
|
|
||||||
.unwrap_or_else(|| "qwen2.5-72b-instruct:free".to_string())
|
|
||||||
} else if let Some(def) = registry.find(&canonical_backend) {
|
|
||||||
std::env::var(&def.model_env)
|
|
||||||
.ok()
|
|
||||||
.or_else(|| settings.selected_model.clone())
|
|
||||||
.unwrap_or_else(|| def.default_model.clone())
|
|
||||||
} else {
|
|
||||||
settings
|
|
||||||
.selected_model
|
|
||||||
.clone()
|
|
||||||
.unwrap_or_else(|| "unknown".to_string())
|
|
||||||
};
|
|
||||||
|
|
||||||
(canonical_backend, model)
|
|
||||||
}
|
|
||||||
|
|
||||||
fn load_settings(config_path: Option<&Path>) -> Settings {
|
|
||||||
if let Some(path) = config_path {
|
|
||||||
Settings::load_toml(path).ok().flatten().unwrap_or_default()
|
|
||||||
} else {
|
|
||||||
let toml_path = config_toml_path();
|
|
||||||
if toml_path.exists() {
|
|
||||||
Settings::load_toml(&toml_path)
|
|
||||||
.ok()
|
|
||||||
.flatten()
|
|
||||||
.unwrap_or_default()
|
|
||||||
} else {
|
|
||||||
Settings::load()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn save_settings(settings: &Settings, config_path: Option<&Path>) -> anyhow::Result<()> {
|
|
||||||
let path = config_path
|
|
||||||
.map(|p| p.to_path_buf())
|
|
||||||
.unwrap_or_else(config_toml_path);
|
|
||||||
|
|
||||||
settings
|
|
||||||
.save_toml(&path)
|
|
||||||
.map_err(|e| anyhow::anyhow!("{}", e))?;
|
|
||||||
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
|
|
||||||
fn config_toml_path() -> std::path::PathBuf {
|
|
||||||
crate::bootstrap::ironclaw_base_dir().join("config.toml")
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Try to fetch the live model list from a provider.
|
|
||||||
///
|
|
||||||
/// Best-effort: returns `None` if config loading, provider creation, or the
|
|
||||||
/// `list_models()` call fails (missing API key, network error, etc.).
|
|
||||||
async fn try_fetch_models(provider_id: &str, config_path: Option<&Path>) -> Option<Vec<String>> {
|
|
||||||
let config = crate::config::Config::from_env_with_toml(config_path)
|
|
||||||
.await
|
|
||||||
.ok()?;
|
|
||||||
|
|
||||||
// Override backend to the requested provider so create_llm_provider
|
|
||||||
// constructs the right one.
|
|
||||||
let mut llm_config = config.llm.clone();
|
|
||||||
llm_config.backend = provider_id.to_string();
|
|
||||||
|
|
||||||
// For registry providers, resolve the RegistryProviderConfig if not
|
|
||||||
// already set for this backend.
|
|
||||||
if provider_id != "nearai" && provider_id != "bedrock" {
|
|
||||||
let registry = ProviderRegistry::load();
|
|
||||||
if let Some(def) = registry.find(provider_id)
|
|
||||||
&& llm_config
|
|
||||||
.provider
|
|
||||||
.as_ref()
|
|
||||||
.is_none_or(|p| p.provider_id != def.id)
|
|
||||||
{
|
|
||||||
// Build a minimal RegistryProviderConfig from env + registry
|
|
||||||
let api_key = def
|
|
||||||
.api_key_env
|
|
||||||
.as_ref()
|
|
||||||
.and_then(|env| std::env::var(env).ok());
|
|
||||||
if def.api_key_required && api_key.is_none() {
|
|
||||||
return None;
|
|
||||||
}
|
|
||||||
let base_url = def.default_base_url.clone().unwrap_or_default();
|
|
||||||
llm_config.provider = Some(crate::llm::RegistryProviderConfig {
|
|
||||||
protocol: def.protocol,
|
|
||||||
provider_id: def.id.clone(),
|
|
||||||
model: def.default_model.clone(),
|
|
||||||
api_key: api_key.map(secrecy::SecretString::from),
|
|
||||||
base_url,
|
|
||||||
extra_headers: Vec::new(),
|
|
||||||
oauth_token: None,
|
|
||||||
is_codex_chatgpt: false,
|
|
||||||
refresh_token: None,
|
|
||||||
auth_path: None,
|
|
||||||
cache_retention: Default::default(),
|
|
||||||
unsupported_params: def.unsupported_params.clone(),
|
|
||||||
});
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
let session = crate::llm::create_session_manager(config.llm.session.clone()).await;
|
|
||||||
let provider = crate::llm::create_llm_provider(&llm_config, session)
|
|
||||||
.await
|
|
||||||
.ok()?;
|
|
||||||
provider.list_models().await.ok().filter(|m| !m.is_empty())
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Print available models section (text output).
|
|
||||||
fn print_model_list(models: &Option<Vec<String>>, active_model: Option<&String>) {
|
|
||||||
match models {
|
|
||||||
Some(models) => {
|
|
||||||
println!("\n Available models ({}):", models.len());
|
|
||||||
for m in models {
|
|
||||||
let marker = active_model
|
|
||||||
.filter(|a| a.as_str() == m)
|
|
||||||
.map(|_| " (active)")
|
|
||||||
.unwrap_or("");
|
|
||||||
println!(" {}{}", m, marker);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
None => {
|
|
||||||
println!(
|
|
||||||
"\n Could not fetch model list (missing credentials or provider unavailable)."
|
|
||||||
);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Also update `~/.ironclaw/.env` so changes take effect immediately.
|
|
||||||
///
|
|
||||||
/// Skipped when `config_path` is `Some` (custom `--config`), because the user
|
|
||||||
/// is explicitly targeting a different config file and we must not pollute the
|
|
||||||
/// default profile's `.env`.
|
|
||||||
fn sync_to_dotenv(config_path: Option<&Path>, vars: &[(&str, &str)]) {
|
|
||||||
if config_path.is_some() {
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
if let Err(e) = crate::bootstrap::upsert_bootstrap_vars(vars) {
|
|
||||||
eprintln!("Warning: failed to update .env: {}", e);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ─── status ───────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
fn cmd_status(json: bool, config_path: Option<&Path>) -> anyhow::Result<()> {
|
|
||||||
let settings = load_settings(config_path);
|
|
||||||
let (backend, model) = resolve_active_from_settings(&settings);
|
|
||||||
let registry = ProviderRegistry::load();
|
|
||||||
|
|
||||||
let fallback = std::env::var("NEARAI_FALLBACK_MODEL").ok();
|
|
||||||
let cheap = std::env::var("NEARAI_CHEAP_MODEL").ok();
|
|
||||||
|
|
||||||
let description = if backend == "nearai" {
|
|
||||||
"NEAR AI inference (default)".to_string()
|
|
||||||
} else {
|
|
||||||
registry
|
|
||||||
.find(&backend)
|
|
||||||
.map(|d| d.description.clone())
|
|
||||||
.unwrap_or_default()
|
|
||||||
};
|
|
||||||
|
|
||||||
if json {
|
|
||||||
let v = serde_json::json!({
|
|
||||||
"provider": backend,
|
|
||||||
"model": model,
|
|
||||||
"description": description,
|
|
||||||
"fallback_model": fallback,
|
|
||||||
"cheap_model": cheap,
|
|
||||||
});
|
|
||||||
println!(
|
|
||||||
"{}",
|
|
||||||
serde_json::to_string_pretty(&v).unwrap_or_else(|_| "{}".to_string())
|
|
||||||
);
|
|
||||||
return Ok(());
|
|
||||||
}
|
|
||||||
|
|
||||||
println!("Provider: {} ({})", backend, description);
|
|
||||||
println!("Model: {}", model);
|
|
||||||
if let Some(ref fb) = fallback {
|
|
||||||
println!("Fallback: {}", fb);
|
|
||||||
}
|
|
||||||
if let Some(ref ch) = cheap {
|
|
||||||
println!("Cheap: {}", ch);
|
|
||||||
}
|
|
||||||
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
|
|
||||||
// ─── set ──────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
fn cmd_set_model(model: &str, config_path: Option<&Path>) -> anyhow::Result<()> {
|
|
||||||
let trimmed = model.trim();
|
|
||||||
if trimmed.is_empty() {
|
|
||||||
anyhow::bail!("Model name cannot be empty");
|
|
||||||
}
|
|
||||||
|
|
||||||
let mut settings = load_settings(config_path);
|
|
||||||
let registry = ProviderRegistry::load();
|
|
||||||
|
|
||||||
// Warn if model name doesn't match any known provider's default model
|
|
||||||
let known_model = registry.all().iter().any(|d| d.default_model == trimmed)
|
|
||||||
|| trimmed.contains("qwen") // nearai models
|
|
||||||
|| trimmed.contains("llama")
|
|
||||||
|| trimmed.contains("gpt")
|
|
||||||
|| trimmed.contains("claude")
|
|
||||||
|| trimmed.contains("gemini")
|
|
||||||
|| trimmed.contains("mistral");
|
|
||||||
if !known_model {
|
|
||||||
eprintln!(
|
|
||||||
"Warning: '{}' is not a recognized model name. Proceeding anyway.",
|
|
||||||
trimmed
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
settings.selected_model = Some(trimmed.to_string());
|
|
||||||
save_settings(&settings, config_path)?;
|
|
||||||
|
|
||||||
let backend = std::env::var("LLM_BACKEND")
|
|
||||||
.ok()
|
|
||||||
.or_else(|| settings.llm_backend.clone())
|
|
||||||
.unwrap_or_else(|| "nearai".to_string());
|
|
||||||
|
|
||||||
// Also write to .env so the change takes effect immediately
|
|
||||||
let model_env = if backend == "nearai" {
|
|
||||||
"NEARAI_MODEL".to_string()
|
|
||||||
} else {
|
|
||||||
registry
|
|
||||||
.find(&backend)
|
|
||||||
.map(|d| d.model_env.clone())
|
|
||||||
.unwrap_or_default()
|
|
||||||
};
|
|
||||||
if !model_env.is_empty() {
|
|
||||||
sync_to_dotenv(config_path, &[(&model_env, trimmed)]);
|
|
||||||
}
|
|
||||||
|
|
||||||
println!("Model set to '{}' (provider: {})", trimmed, backend);
|
|
||||||
println!(
|
|
||||||
"Saved to {}",
|
|
||||||
config_path
|
|
||||||
.map(|p| p.display().to_string())
|
|
||||||
.unwrap_or_else(|| config_toml_path().display().to_string())
|
|
||||||
);
|
|
||||||
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
|
|
||||||
// ─── set-provider ─────────────────────────────────────────────────
|
|
||||||
|
|
||||||
fn cmd_set_provider(
|
|
||||||
provider: &str,
|
|
||||||
model: Option<&str>,
|
|
||||||
config_path: Option<&Path>,
|
|
||||||
) -> anyhow::Result<()> {
|
|
||||||
let registry = ProviderRegistry::load();
|
|
||||||
|
|
||||||
// Validate and normalize provider
|
|
||||||
let canonical_id = if provider == "nearai" || provider == "near_ai" || provider == "near" {
|
|
||||||
"nearai".to_string()
|
|
||||||
} else {
|
|
||||||
let def = registry.find(provider).ok_or_else(|| {
|
|
||||||
let known: Vec<&str> = std::iter::once("nearai")
|
|
||||||
.chain(registry.all().iter().map(|d| d.id.as_str()))
|
|
||||||
.collect();
|
|
||||||
anyhow::anyhow!(
|
|
||||||
"Unknown provider '{}'. Known providers: {}",
|
|
||||||
provider,
|
|
||||||
known.join(", ")
|
|
||||||
)
|
|
||||||
})?;
|
|
||||||
def.id.clone()
|
|
||||||
};
|
|
||||||
|
|
||||||
// Resolve model: explicit > provider default
|
|
||||||
let resolved_model = if let Some(m) = model {
|
|
||||||
m.to_string()
|
|
||||||
} else if canonical_id == "nearai" {
|
|
||||||
"qwen2.5-72b-instruct:free".to_string()
|
|
||||||
} else if let Some(def) = registry.find(&canonical_id) {
|
|
||||||
def.default_model.clone()
|
|
||||||
} else {
|
|
||||||
"default".to_string()
|
|
||||||
};
|
|
||||||
|
|
||||||
let mut settings = load_settings(config_path);
|
|
||||||
settings.llm_backend = Some(canonical_id.clone());
|
|
||||||
settings.selected_model = Some(resolved_model.clone());
|
|
||||||
save_settings(&settings, config_path)?;
|
|
||||||
|
|
||||||
// Also write to .env so the change takes effect immediately
|
|
||||||
let model_env = if canonical_id == "nearai" {
|
|
||||||
"NEARAI_MODEL".to_string()
|
|
||||||
} else {
|
|
||||||
registry
|
|
||||||
.find(&canonical_id)
|
|
||||||
.map(|d| d.model_env.clone())
|
|
||||||
.unwrap_or_default()
|
|
||||||
};
|
|
||||||
let mut vars: Vec<(&str, &str)> = vec![("LLM_BACKEND", &canonical_id)];
|
|
||||||
if !model_env.is_empty() {
|
|
||||||
vars.push((&model_env, &resolved_model));
|
|
||||||
}
|
|
||||||
sync_to_dotenv(config_path, &vars);
|
|
||||||
|
|
||||||
println!(
|
|
||||||
"Provider set to '{}', model set to '{}'",
|
|
||||||
canonical_id, resolved_model
|
|
||||||
);
|
|
||||||
println!(
|
|
||||||
"Saved to {}",
|
|
||||||
config_path
|
|
||||||
.map(|p| p.display().to_string())
|
|
||||||
.unwrap_or_else(|| config_toml_path().display().to_string())
|
|
||||||
);
|
|
||||||
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
|
|
||||||
// ─── list ─────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
/// List all providers with their default models.
|
|
||||||
async fn cmd_list_providers(
|
|
||||||
verbose: bool,
|
|
||||||
json: bool,
|
|
||||||
config_path: Option<&Path>,
|
|
||||||
) -> anyhow::Result<()> {
|
|
||||||
let registry = ProviderRegistry::load();
|
|
||||||
let (active_backend, active_model) = resolve_active(config_path);
|
|
||||||
|
|
||||||
if json {
|
|
||||||
let mut entries: Vec<serde_json::Value> = Vec::new();
|
|
||||||
|
|
||||||
// NEAR AI (not in registry)
|
|
||||||
let nearai_active = active_backend == "nearai";
|
|
||||||
entries.push(serde_json::json!({
|
|
||||||
"id": "nearai",
|
|
||||||
"description": "NEAR AI inference (default)",
|
|
||||||
"default_model": "qwen2.5-72b-instruct:free",
|
|
||||||
"active": nearai_active,
|
|
||||||
"active_model": if nearai_active { Some(&active_model) } else { None },
|
|
||||||
}));
|
|
||||||
|
|
||||||
for def in registry.all() {
|
|
||||||
let is_active = active_backend == def.id;
|
|
||||||
let mut v = serde_json::json!({
|
|
||||||
"id": def.id,
|
|
||||||
"description": def.description,
|
|
||||||
"default_model": def.default_model,
|
|
||||||
"protocol": format!("{:?}", def.protocol),
|
|
||||||
"active": is_active,
|
|
||||||
});
|
|
||||||
if is_active {
|
|
||||||
v["active_model"] = serde_json::json!(active_model);
|
|
||||||
}
|
|
||||||
if verbose {
|
|
||||||
v["aliases"] = serde_json::json!(def.aliases);
|
|
||||||
v["model_env"] = serde_json::json!(def.model_env);
|
|
||||||
v["api_key_env"] = serde_json::json!(def.api_key_env);
|
|
||||||
v["api_key_required"] = serde_json::json!(def.api_key_required);
|
|
||||||
if let Some(ref url) = def.default_base_url {
|
|
||||||
v["base_url"] = serde_json::json!(url);
|
|
||||||
}
|
|
||||||
if let Some(ref setup) = def.setup {
|
|
||||||
v["can_list_models"] = serde_json::json!(setup.can_list_models());
|
|
||||||
}
|
|
||||||
}
|
|
||||||
entries.push(v);
|
|
||||||
}
|
|
||||||
|
|
||||||
println!(
|
|
||||||
"{}",
|
|
||||||
serde_json::to_string_pretty(&entries).unwrap_or_else(|_| "[]".to_string())
|
|
||||||
);
|
|
||||||
return Ok(());
|
|
||||||
}
|
|
||||||
|
|
||||||
let providers = registry.all();
|
|
||||||
|
|
||||||
println!("Active: {} (model: {})\n", active_backend, active_model);
|
|
||||||
println!(
|
|
||||||
"{} provider(s) available:\n",
|
|
||||||
providers.len() + 1 // +1 for NEAR AI
|
|
||||||
);
|
|
||||||
|
|
||||||
// NEAR AI (not in registry)
|
|
||||||
let nearai_marker = if active_backend == "nearai" { " *" } else { "" };
|
|
||||||
if verbose {
|
|
||||||
println!(" nearai{}", nearai_marker);
|
|
||||||
println!(" Description: NEAR AI inference (default)");
|
|
||||||
println!(" Default model: qwen2.5-72b-instruct:free");
|
|
||||||
println!(" Model env: NEARAI_MODEL");
|
|
||||||
if active_backend == "nearai" {
|
|
||||||
println!(" Active model: {}", active_model);
|
|
||||||
}
|
|
||||||
println!();
|
|
||||||
} else {
|
|
||||||
println!(
|
|
||||||
" {:<22} {:<40} NEAR AI inference (default)",
|
|
||||||
format!("nearai{nearai_marker}"),
|
|
||||||
"qwen2.5-72b-instruct:free"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
for def in providers {
|
|
||||||
let is_active = active_backend == def.id;
|
|
||||||
let marker = if is_active { " *" } else { "" };
|
|
||||||
|
|
||||||
if verbose {
|
|
||||||
println!(" {}{}", def.id, marker);
|
|
||||||
println!(" Description: {}", def.description);
|
|
||||||
println!(" Default model: {}", def.default_model);
|
|
||||||
println!(" Protocol: {:?}", def.protocol);
|
|
||||||
println!(" Model env: {}", def.model_env);
|
|
||||||
if let Some(ref env) = def.api_key_env {
|
|
||||||
println!(
|
|
||||||
" API key env: {} ({})",
|
|
||||||
env,
|
|
||||||
if def.api_key_required {
|
|
||||||
"required"
|
|
||||||
} else {
|
|
||||||
"optional"
|
|
||||||
}
|
|
||||||
);
|
|
||||||
}
|
|
||||||
if let Some(ref url) = def.default_base_url {
|
|
||||||
println!(" Base URL: {}", url);
|
|
||||||
}
|
|
||||||
if !def.aliases.is_empty() {
|
|
||||||
println!(" Aliases: {}", def.aliases.join(", "));
|
|
||||||
}
|
|
||||||
if is_active {
|
|
||||||
println!(" Active model: {}", active_model);
|
|
||||||
}
|
|
||||||
println!();
|
|
||||||
} else {
|
|
||||||
let model_display = if is_active {
|
|
||||||
active_model.clone()
|
|
||||||
} else {
|
|
||||||
def.default_model.clone()
|
|
||||||
};
|
|
||||||
println!(
|
|
||||||
" {:<22} {:<40} {}",
|
|
||||||
format!("{}{marker}", def.id),
|
|
||||||
model_display,
|
|
||||||
def.description,
|
|
||||||
);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if !verbose {
|
|
||||||
println!();
|
|
||||||
println!("* = active provider. Use --verbose for details.");
|
|
||||||
}
|
|
||||||
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Show details for a specific provider.
|
|
||||||
async fn cmd_show_provider(
|
|
||||||
id: &str,
|
|
||||||
verbose: bool,
|
|
||||||
json: bool,
|
|
||||||
config_path: Option<&Path>,
|
|
||||||
) -> anyhow::Result<()> {
|
|
||||||
let registry = ProviderRegistry::load();
|
|
||||||
let (active_backend, active_model) = resolve_active(config_path);
|
|
||||||
|
|
||||||
// Resolve canonical ID for model fetching
|
|
||||||
let canonical_id = if id == "nearai" || id == "near_ai" || id == "near" {
|
|
||||||
"nearai".to_string()
|
|
||||||
} else {
|
|
||||||
registry
|
|
||||||
.find(id)
|
|
||||||
.map(|d| d.id.clone())
|
|
||||||
.unwrap_or_else(|| id.to_string())
|
|
||||||
};
|
|
||||||
|
|
||||||
// Try to fetch live model list from the provider
|
|
||||||
let live_models = try_fetch_models(&canonical_id, config_path).await;
|
|
||||||
|
|
||||||
// Check NEAR AI first (not in registry)
|
|
||||||
if id == "nearai" || id == "near_ai" || id == "near" {
|
|
||||||
let is_active = active_backend == "nearai";
|
|
||||||
if json {
|
|
||||||
let mut v = serde_json::json!({
|
|
||||||
"id": "nearai",
|
|
||||||
"description": "NEAR AI inference (default)",
|
|
||||||
"default_model": "qwen2.5-72b-instruct:free",
|
|
||||||
"model_env": "NEARAI_MODEL",
|
|
||||||
"active": is_active,
|
|
||||||
});
|
|
||||||
if is_active {
|
|
||||||
v["active_model"] = serde_json::json!(active_model);
|
|
||||||
}
|
|
||||||
if let Some(ref models) = live_models {
|
|
||||||
v["available_models"] = serde_json::json!(models);
|
|
||||||
}
|
|
||||||
println!(
|
|
||||||
"{}",
|
|
||||||
serde_json::to_string_pretty(&v).unwrap_or_else(|_| "{}".to_string())
|
|
||||||
);
|
|
||||||
} else {
|
|
||||||
println!("Provider: nearai");
|
|
||||||
println!(" Description: NEAR AI inference (default)");
|
|
||||||
println!(" Default model: qwen2.5-72b-instruct:free");
|
|
||||||
println!(" Model env: NEARAI_MODEL");
|
|
||||||
println!(" Active: {}", if is_active { "yes" } else { "no" });
|
|
||||||
if is_active {
|
|
||||||
println!(" Active model: {}", active_model);
|
|
||||||
}
|
|
||||||
print_model_list(&live_models, is_active.then_some(&active_model));
|
|
||||||
}
|
|
||||||
return Ok(());
|
|
||||||
}
|
|
||||||
|
|
||||||
let def = registry.find(id).ok_or_else(|| {
|
|
||||||
let known: Vec<&str> = std::iter::once("nearai")
|
|
||||||
.chain(registry.all().iter().map(|d| d.id.as_str()))
|
|
||||||
.collect();
|
|
||||||
anyhow::anyhow!(
|
|
||||||
"Unknown provider '{}'. Known providers: {}",
|
|
||||||
id,
|
|
||||||
known.join(", ")
|
|
||||||
)
|
|
||||||
})?;
|
|
||||||
|
|
||||||
let is_active = active_backend == def.id;
|
|
||||||
|
|
||||||
if json {
|
|
||||||
let mut v = serde_json::json!({
|
|
||||||
"id": def.id,
|
|
||||||
"description": def.description,
|
|
||||||
"protocol": format!("{:?}", def.protocol),
|
|
||||||
"default_model": def.default_model,
|
|
||||||
"model_env": def.model_env,
|
|
||||||
"api_key_env": def.api_key_env,
|
|
||||||
"api_key_required": def.api_key_required,
|
|
||||||
"aliases": def.aliases,
|
|
||||||
"active": is_active,
|
|
||||||
});
|
|
||||||
if let Some(ref url) = def.default_base_url {
|
|
||||||
v["base_url"] = serde_json::json!(url);
|
|
||||||
}
|
|
||||||
if let Some(ref setup) = def.setup {
|
|
||||||
v["can_list_models"] = serde_json::json!(setup.can_list_models());
|
|
||||||
v["display_name"] = serde_json::json!(setup.display_name());
|
|
||||||
}
|
|
||||||
if is_active {
|
|
||||||
v["active_model"] = serde_json::json!(active_model);
|
|
||||||
}
|
|
||||||
if verbose && !def.unsupported_params.is_empty() {
|
|
||||||
v["unsupported_params"] = serde_json::json!(def.unsupported_params);
|
|
||||||
}
|
|
||||||
if let Some(ref models) = live_models {
|
|
||||||
v["available_models"] = serde_json::json!(models);
|
|
||||||
}
|
|
||||||
println!(
|
|
||||||
"{}",
|
|
||||||
serde_json::to_string_pretty(&v).unwrap_or_else(|_| "{}".to_string())
|
|
||||||
);
|
|
||||||
return Ok(());
|
|
||||||
}
|
|
||||||
|
|
||||||
println!("Provider: {}", def.id);
|
|
||||||
println!(" Description: {}", def.description);
|
|
||||||
println!(" Protocol: {:?}", def.protocol);
|
|
||||||
println!(" Default model: {}", def.default_model);
|
|
||||||
println!(" Model env: {}", def.model_env);
|
|
||||||
if let Some(ref env) = def.api_key_env {
|
|
||||||
println!(
|
|
||||||
" API key env: {} ({})",
|
|
||||||
env,
|
|
||||||
if def.api_key_required {
|
|
||||||
"required"
|
|
||||||
} else {
|
|
||||||
"optional"
|
|
||||||
}
|
|
||||||
);
|
|
||||||
}
|
|
||||||
if let Some(ref url) = def.default_base_url {
|
|
||||||
println!(" Base URL: {}", url);
|
|
||||||
}
|
|
||||||
if !def.aliases.is_empty() {
|
|
||||||
println!(" Aliases: {}", def.aliases.join(", "));
|
|
||||||
}
|
|
||||||
if let Some(ref setup) = def.setup {
|
|
||||||
println!(
|
|
||||||
" List models: {}",
|
|
||||||
if setup.can_list_models() {
|
|
||||||
"supported"
|
|
||||||
} else {
|
|
||||||
"not supported"
|
|
||||||
}
|
|
||||||
);
|
|
||||||
println!(" Display name: {}", setup.display_name());
|
|
||||||
}
|
|
||||||
if !def.unsupported_params.is_empty() {
|
|
||||||
println!(" Unsupported: {}", def.unsupported_params.join(", "));
|
|
||||||
}
|
|
||||||
println!(" Active: {}", if is_active { "yes" } else { "no" });
|
|
||||||
if is_active {
|
|
||||||
println!(" Active model: {}", active_model);
|
|
||||||
}
|
|
||||||
print_model_list(&live_models, is_active.then_some(&active_model));
|
|
||||||
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn resolve_active_defaults_to_nearai() {
|
|
||||||
let settings = Settings::default();
|
|
||||||
assert!(settings.llm_backend.is_none());
|
|
||||||
assert!(settings.selected_model.is_none());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn registry_loads_all_providers() {
|
|
||||||
let registry = ProviderRegistry::load();
|
|
||||||
let all = registry.all();
|
|
||||||
assert!(
|
|
||||||
all.len() >= 10,
|
|
||||||
"should have at least 10 built-in providers, got {}",
|
|
||||||
all.len()
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn registry_find_by_alias() {
|
|
||||||
let registry = ProviderRegistry::load();
|
|
||||||
let def = registry
|
|
||||||
.find("claude")
|
|
||||||
.expect("claude alias should resolve");
|
|
||||||
assert_eq!(def.id, "anthropic");
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn all_providers_have_description() {
|
|
||||||
let registry = ProviderRegistry::load();
|
|
||||||
for def in registry.all() {
|
|
||||||
assert!(
|
|
||||||
!def.description.is_empty(),
|
|
||||||
"provider {} should have a description",
|
|
||||||
def.id
|
|
||||||
);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn set_model_persists_to_toml() {
|
|
||||||
let dir = tempfile::tempdir().expect("create temp dir");
|
|
||||||
let toml_path = dir.path().join("config.toml");
|
|
||||||
|
|
||||||
cmd_set_model("gpt-5-mini", Some(&toml_path)).expect("set model");
|
|
||||||
|
|
||||||
let settings = Settings::load_toml(&toml_path)
|
|
||||||
.expect("read toml")
|
|
||||||
.expect("should have settings");
|
|
||||||
assert_eq!(settings.selected_model.as_deref(), Some("gpt-5-mini"));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn set_provider_validates_unknown() {
|
|
||||||
let dir = tempfile::tempdir().expect("create temp dir");
|
|
||||||
let toml_path = dir.path().join("config.toml");
|
|
||||||
|
|
||||||
let result = cmd_set_provider("nonexistent_provider", None, Some(&toml_path));
|
|
||||||
assert!(result.is_err());
|
|
||||||
let err = result.unwrap_err().to_string();
|
|
||||||
assert!(
|
|
||||||
err.contains("Unknown provider"),
|
|
||||||
"should mention unknown provider: {}",
|
|
||||||
err
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn set_provider_persists_to_toml() {
|
|
||||||
let dir = tempfile::tempdir().expect("create temp dir");
|
|
||||||
let toml_path = dir.path().join("config.toml");
|
|
||||||
|
|
||||||
cmd_set_provider("groq", None, Some(&toml_path)).expect("set provider");
|
|
||||||
|
|
||||||
let settings = Settings::load_toml(&toml_path)
|
|
||||||
.expect("read toml")
|
|
||||||
.expect("should have settings");
|
|
||||||
assert_eq!(settings.llm_backend.as_deref(), Some("groq"));
|
|
||||||
assert_eq!(
|
|
||||||
settings.selected_model.as_deref(),
|
|
||||||
Some("llama-3.3-70b-versatile")
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn set_provider_with_custom_model() {
|
|
||||||
let dir = tempfile::tempdir().expect("create temp dir");
|
|
||||||
let toml_path = dir.path().join("config.toml");
|
|
||||||
|
|
||||||
cmd_set_provider("anthropic", Some("claude-opus-4-6"), Some(&toml_path))
|
|
||||||
.expect("set provider with model");
|
|
||||||
|
|
||||||
let settings = Settings::load_toml(&toml_path)
|
|
||||||
.expect("read toml")
|
|
||||||
.expect("should have settings");
|
|
||||||
assert_eq!(settings.llm_backend.as_deref(), Some("anthropic"));
|
|
||||||
assert_eq!(settings.selected_model.as_deref(), Some("claude-opus-4-6"));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn custom_config_does_not_pollute_default_dotenv() {
|
|
||||||
let dir = tempfile::tempdir().expect("create temp dir");
|
|
||||||
let toml_path = dir.path().join("config.toml");
|
|
||||||
|
|
||||||
// With a custom config path, sync_to_dotenv should be a no-op
|
|
||||||
// (it returns early when config_path is Some).
|
|
||||||
// We verify by checking that cmd_set_provider succeeds without
|
|
||||||
// trying to write to the default ~/.ironclaw/.env.
|
|
||||||
cmd_set_provider("groq", None, Some(&toml_path)).expect("set provider with custom config");
|
|
||||||
|
|
||||||
let settings = Settings::load_toml(&toml_path)
|
|
||||||
.expect("read toml")
|
|
||||||
.expect("should have settings");
|
|
||||||
assert_eq!(settings.llm_backend.as_deref(), Some("groq"));
|
|
||||||
// The key assertion is that no error was thrown trying to write
|
|
||||||
// to the default .env — sync_to_dotenv skipped it.
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn set_model_rejects_empty_name() {
|
|
||||||
let dir = tempfile::tempdir().expect("create temp dir");
|
|
||||||
let toml_path = dir.path().join("config.toml");
|
|
||||||
|
|
||||||
let result = cmd_set_model("", Some(&toml_path));
|
|
||||||
assert!(result.is_err());
|
|
||||||
assert!(
|
|
||||||
result.unwrap_err().to_string().contains("cannot be empty"),
|
|
||||||
"should reject empty model name"
|
|
||||||
);
|
|
||||||
|
|
||||||
let result2 = cmd_set_model(" ", Some(&toml_path));
|
|
||||||
assert!(result2.is_err());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn set_provider_normalizes_alias() {
|
|
||||||
let dir = tempfile::tempdir().expect("create temp dir");
|
|
||||||
let toml_path = dir.path().join("config.toml");
|
|
||||||
|
|
||||||
cmd_set_provider("claude", None, Some(&toml_path)).expect("set via alias");
|
|
||||||
|
|
||||||
let settings = Settings::load_toml(&toml_path)
|
|
||||||
.expect("read toml")
|
|
||||||
.expect("should have settings");
|
|
||||||
assert_eq!(
|
|
||||||
settings.llm_backend.as_deref(),
|
|
||||||
Some("anthropic"),
|
|
||||||
"alias should be normalized to canonical ID"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -447,8 +447,8 @@ pub struct PendingOAuthFlow {
|
|||||||
pub user_id: String,
|
pub user_id: String,
|
||||||
/// Secrets store reference for token persistence.
|
/// Secrets store reference for token persistence.
|
||||||
pub secrets: Arc<dyn SecretsStore + Send + Sync>,
|
pub secrets: Arc<dyn SecretsStore + Send + Sync>,
|
||||||
/// SSE broadcast manager for notifying the web UI.
|
/// SSE broadcast sender for notifying the web UI.
|
||||||
pub sse_manager: Option<Arc<crate::channels::web::sse::SseManager>>,
|
pub sse_sender: Option<tokio::sync::broadcast::Sender<crate::channels::web::types::SseEvent>>,
|
||||||
/// Gateway auth token for authenticating with the platform token exchange proxy.
|
/// Gateway auth token for authenticating with the platform token exchange proxy.
|
||||||
pub gateway_token: Option<String>,
|
pub gateway_token: Option<String>,
|
||||||
/// Additional form params for the token exchange request.
|
/// Additional form params for the token exchange request.
|
||||||
|
|||||||
@@ -20,7 +20,6 @@ Commands:
|
|||||||
service Manage OS service
|
service Manage OS service
|
||||||
skills Manage skills
|
skills Manage skills
|
||||||
hooks Manage lifecycle hooks
|
hooks Manage lifecycle hooks
|
||||||
models Manage LLM providers and models
|
|
||||||
doctor Run diagnostics
|
doctor Run diagnostics
|
||||||
logs View and manage gateway logs
|
logs View and manage gateway logs
|
||||||
status Show system status
|
status Show system status
|
||||||
|
|||||||
@@ -20,7 +20,6 @@ Commands:
|
|||||||
service Manage OS service
|
service Manage OS service
|
||||||
skills Manage skills
|
skills Manage skills
|
||||||
hooks Manage lifecycle hooks
|
hooks Manage lifecycle hooks
|
||||||
models Manage LLM providers and models
|
|
||||||
doctor Run diagnostics
|
doctor Run diagnostics
|
||||||
logs View and manage gateway logs
|
logs View and manage gateway logs
|
||||||
status Show system status
|
status Show system status
|
||||||
|
|||||||
@@ -23,7 +23,6 @@ Commands:
|
|||||||
service Manage OS service
|
service Manage OS service
|
||||||
skills Manage skills
|
skills Manage skills
|
||||||
hooks Manage lifecycle hooks
|
hooks Manage lifecycle hooks
|
||||||
models Manage LLM providers and models
|
|
||||||
doctor Run diagnostics
|
doctor Run diagnostics
|
||||||
logs View and manage gateway logs
|
logs View and manage gateway logs
|
||||||
status Show system status
|
status Show system status
|
||||||
|
|||||||
@@ -23,7 +23,6 @@ Commands:
|
|||||||
service Manage OS service
|
service Manage OS service
|
||||||
skills Manage skills
|
skills Manage skills
|
||||||
hooks Manage lifecycle hooks
|
hooks Manage lifecycle hooks
|
||||||
models Manage LLM providers and models
|
|
||||||
doctor Run diagnostics
|
doctor Run diagnostics
|
||||||
logs View and manage gateway logs
|
logs View and manage gateway logs
|
||||||
status Show system status
|
status Show system status
|
||||||
|
|||||||
@@ -2,7 +2,6 @@ use std::collections::HashMap;
|
|||||||
use std::path::PathBuf;
|
use std::path::PathBuf;
|
||||||
|
|
||||||
use secrecy::SecretString;
|
use secrecy::SecretString;
|
||||||
use serde::Deserialize;
|
|
||||||
|
|
||||||
use crate::bootstrap::ironclaw_base_dir;
|
use crate::bootstrap::ironclaw_base_dir;
|
||||||
use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env};
|
use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env};
|
||||||
@@ -46,26 +45,6 @@ pub struct GatewayConfig {
|
|||||||
/// Bearer token for authentication. Random hex generated at startup if unset.
|
/// Bearer token for authentication. Random hex generated at startup if unset.
|
||||||
pub auth_token: Option<String>,
|
pub auth_token: Option<String>,
|
||||||
pub user_id: String,
|
pub user_id: String,
|
||||||
/// Additional user scopes for workspace reads.
|
|
||||||
///
|
|
||||||
/// When set, the workspace will be able to read (search, read, list) from
|
|
||||||
/// these additional user scopes while writes remain isolated to `user_id`.
|
|
||||||
/// Parsed from `WORKSPACE_READ_SCOPES` (comma-separated).
|
|
||||||
pub workspace_read_scopes: Vec<String>,
|
|
||||||
/// Memory layer definitions (JSON in env var, or from external config).
|
|
||||||
pub memory_layers: Vec<crate::workspace::layer::MemoryLayer>,
|
|
||||||
/// Multi-user token map. When set, each token maps to a user identity.
|
|
||||||
/// Parsed from `GATEWAY_USER_TOKENS` (JSON string). When absent, falls back
|
|
||||||
/// to single-user mode via `auth_token` + `user_id`.
|
|
||||||
pub user_tokens: Option<HashMap<String, UserTokenConfig>>,
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Per-user token configuration for multi-user mode.
|
|
||||||
#[derive(Debug, Clone, Deserialize)]
|
|
||||||
pub struct UserTokenConfig {
|
|
||||||
pub user_id: String,
|
|
||||||
#[serde(default)]
|
|
||||||
pub workspace_read_scopes: Vec<String>,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Signal channel configuration (signal-cli daemon HTTP/JSON-RPC).
|
/// Signal channel configuration (signal-cli daemon HTTP/JSON-RPC).
|
||||||
@@ -136,118 +115,6 @@ impl ChannelsConfig {
|
|||||||
.or_else(|| cs.gateway_user_id.clone())
|
.or_else(|| cs.gateway_user_id.clone())
|
||||||
.unwrap_or_else(|| owner_id.to_string());
|
.unwrap_or_else(|| owner_id.to_string());
|
||||||
|
|
||||||
let memory_layers: Vec<crate::workspace::layer::MemoryLayer> =
|
|
||||||
match optional_env("MEMORY_LAYERS")? {
|
|
||||||
Some(json_str) => {
|
|
||||||
serde_json::from_str(&json_str).map_err(|e| ConfigError::InvalidValue {
|
|
||||||
key: "MEMORY_LAYERS".to_string(),
|
|
||||||
message: format!("must be valid JSON array of layer objects: {e}"),
|
|
||||||
})?
|
|
||||||
}
|
|
||||||
None => crate::workspace::layer::MemoryLayer::default_for_user(&user_id),
|
|
||||||
};
|
|
||||||
|
|
||||||
// Validate layer names and scopes
|
|
||||||
for layer in &memory_layers {
|
|
||||||
if layer.name.trim().is_empty() {
|
|
||||||
return Err(ConfigError::InvalidValue {
|
|
||||||
key: "MEMORY_LAYERS".to_string(),
|
|
||||||
message: "layer name must not be empty".to_string(),
|
|
||||||
});
|
|
||||||
}
|
|
||||||
if layer.name.len() > 64 {
|
|
||||||
return Err(ConfigError::InvalidValue {
|
|
||||||
key: "MEMORY_LAYERS".to_string(),
|
|
||||||
message: format!("layer name '{}' exceeds 64 characters", layer.name),
|
|
||||||
});
|
|
||||||
}
|
|
||||||
if !layer
|
|
||||||
.name
|
|
||||||
.chars()
|
|
||||||
.all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-')
|
|
||||||
{
|
|
||||||
return Err(ConfigError::InvalidValue {
|
|
||||||
key: "MEMORY_LAYERS".to_string(),
|
|
||||||
message: format!(
|
|
||||||
"layer name '{}' contains invalid characters \
|
|
||||||
(allowed: a-z, A-Z, 0-9, _, -)",
|
|
||||||
layer.name
|
|
||||||
),
|
|
||||||
});
|
|
||||||
}
|
|
||||||
if layer.scope.trim().is_empty() {
|
|
||||||
return Err(ConfigError::InvalidValue {
|
|
||||||
key: "MEMORY_LAYERS".to_string(),
|
|
||||||
message: format!("layer '{}' has an empty scope", layer.name),
|
|
||||||
});
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Check for duplicate layer names
|
|
||||||
{
|
|
||||||
let mut seen = std::collections::HashSet::new();
|
|
||||||
for layer in &memory_layers {
|
|
||||||
if !seen.insert(&layer.name) {
|
|
||||||
return Err(ConfigError::InvalidValue {
|
|
||||||
key: "MEMORY_LAYERS".to_string(),
|
|
||||||
message: format!("duplicate layer name '{}'", layer.name),
|
|
||||||
});
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
let user_tokens: Option<HashMap<String, UserTokenConfig>> =
|
|
||||||
match optional_env("GATEWAY_USER_TOKENS")? {
|
|
||||||
Some(json_str) => {
|
|
||||||
let tokens: HashMap<String, UserTokenConfig> = serde_json::from_str(
|
|
||||||
&json_str,
|
|
||||||
)
|
|
||||||
.map_err(|e| ConfigError::InvalidValue {
|
|
||||||
key: "GATEWAY_USER_TOKENS".to_string(),
|
|
||||||
message: format!(
|
|
||||||
"must be valid JSON object mapping tokens to user configs: {e}"
|
|
||||||
),
|
|
||||||
})?;
|
|
||||||
if tokens.is_empty() {
|
|
||||||
return Err(ConfigError::InvalidValue {
|
|
||||||
key: "GATEWAY_USER_TOKENS".to_string(),
|
|
||||||
message:
|
|
||||||
"token map is empty — remove the variable to use single-user mode"
|
|
||||||
.to_string(),
|
|
||||||
});
|
|
||||||
}
|
|
||||||
for (tok, cfg) in &tokens {
|
|
||||||
if cfg.user_id.trim().is_empty() {
|
|
||||||
return Err(ConfigError::InvalidValue {
|
|
||||||
key: "GATEWAY_USER_TOKENS".to_string(),
|
|
||||||
message: format!(
|
|
||||||
"token '{}...' has an empty user_id",
|
|
||||||
&tok[..tok.len().min(8)]
|
|
||||||
),
|
|
||||||
});
|
|
||||||
}
|
|
||||||
}
|
|
||||||
Some(tokens)
|
|
||||||
}
|
|
||||||
None => None,
|
|
||||||
};
|
|
||||||
let workspace_read_scopes: Vec<String> = optional_env("WORKSPACE_READ_SCOPES")?
|
|
||||||
.map(|s| {
|
|
||||||
s.split(',')
|
|
||||||
.map(|s| s.trim().to_string())
|
|
||||||
.filter(|s| !s.is_empty())
|
|
||||||
.collect()
|
|
||||||
})
|
|
||||||
.unwrap_or_default();
|
|
||||||
|
|
||||||
for scope in &workspace_read_scopes {
|
|
||||||
if scope.len() > 128 {
|
|
||||||
return Err(ConfigError::InvalidValue {
|
|
||||||
key: "WORKSPACE_READ_SCOPES".to_string(),
|
|
||||||
message: format!("scope '{}...' exceeds 128 characters", &scope[..32]),
|
|
||||||
});
|
|
||||||
}
|
|
||||||
}
|
|
||||||
Some(GatewayConfig {
|
Some(GatewayConfig {
|
||||||
host: optional_env("GATEWAY_HOST")?
|
host: optional_env("GATEWAY_HOST")?
|
||||||
.or_else(|| cs.gateway_host.clone())
|
.or_else(|| cs.gateway_host.clone())
|
||||||
@@ -259,9 +126,6 @@ impl ChannelsConfig {
|
|||||||
auth_token: optional_env("GATEWAY_AUTH_TOKEN")?
|
auth_token: optional_env("GATEWAY_AUTH_TOKEN")?
|
||||||
.or_else(|| cs.gateway_auth_token.clone()),
|
.or_else(|| cs.gateway_auth_token.clone()),
|
||||||
user_id,
|
user_id,
|
||||||
workspace_read_scopes,
|
|
||||||
memory_layers,
|
|
||||||
user_tokens,
|
|
||||||
})
|
})
|
||||||
} else {
|
} else {
|
||||||
None
|
None
|
||||||
@@ -417,9 +281,6 @@ mod tests {
|
|||||||
port: 3000,
|
port: 3000,
|
||||||
auth_token: Some("tok-abc".to_string()),
|
auth_token: Some("tok-abc".to_string()),
|
||||||
user_id: "default".to_string(),
|
user_id: "default".to_string(),
|
||||||
workspace_read_scopes: vec![],
|
|
||||||
memory_layers: vec![],
|
|
||||||
user_tokens: None,
|
|
||||||
};
|
};
|
||||||
assert_eq!(cfg.host, "127.0.0.1");
|
assert_eq!(cfg.host, "127.0.0.1");
|
||||||
assert_eq!(cfg.port, 3000);
|
assert_eq!(cfg.port, 3000);
|
||||||
@@ -434,9 +295,6 @@ mod tests {
|
|||||||
port: 3001,
|
port: 3001,
|
||||||
auth_token: None,
|
auth_token: None,
|
||||||
user_id: "anon".to_string(),
|
user_id: "anon".to_string(),
|
||||||
workspace_read_scopes: vec![],
|
|
||||||
memory_layers: vec![],
|
|
||||||
user_tokens: None,
|
|
||||||
};
|
};
|
||||||
assert!(cfg.auth_token.is_none());
|
assert!(cfg.auth_token.is_none());
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -230,49 +230,6 @@ impl JobStore for LibSqlBackend {
|
|||||||
Ok(jobs)
|
Ok(jobs)
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn list_agent_jobs_for_user(
|
|
||||||
&self,
|
|
||||||
user_id: &str,
|
|
||||||
) -> Result<Vec<AgentJobRecord>, DatabaseError> {
|
|
||||||
let conn = self.connect().await?;
|
|
||||||
let mut rows = conn
|
|
||||||
.query(
|
|
||||||
r#"
|
|
||||||
SELECT id, title, status, user_id, failure_reason,
|
|
||||||
created_at, started_at, completed_at
|
|
||||||
FROM agent_jobs WHERE source = 'direct' AND user_id = ?1
|
|
||||||
ORDER BY created_at DESC
|
|
||||||
"#,
|
|
||||||
params![user_id],
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
|
||||||
|
|
||||||
let mut jobs = Vec::new();
|
|
||||||
while let Some(row) = rows
|
|
||||||
.next()
|
|
||||||
.await
|
|
||||||
.map_err(|e| DatabaseError::Query(e.to_string()))?
|
|
||||||
{
|
|
||||||
let id_str = get_text(&row, 0);
|
|
||||||
let Ok(id) = id_str.parse() else {
|
|
||||||
tracing::warn!("Skipping agent job with invalid UUID: {}", id_str);
|
|
||||||
continue;
|
|
||||||
};
|
|
||||||
jobs.push(AgentJobRecord {
|
|
||||||
id,
|
|
||||||
title: get_text(&row, 1),
|
|
||||||
status: get_text(&row, 2),
|
|
||||||
user_id: get_text(&row, 3),
|
|
||||||
failure_reason: get_opt_text(&row, 4),
|
|
||||||
created_at: get_ts(&row, 5),
|
|
||||||
started_at: get_opt_ts(&row, 6),
|
|
||||||
completed_at: get_opt_ts(&row, 7),
|
|
||||||
});
|
|
||||||
}
|
|
||||||
Ok(jobs)
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn get_agent_job_failure_reason(
|
async fn get_agent_job_failure_reason(
|
||||||
&self,
|
&self,
|
||||||
id: Uuid,
|
id: Uuid,
|
||||||
@@ -320,32 +277,6 @@ impl JobStore for LibSqlBackend {
|
|||||||
Ok(summary)
|
Ok(summary)
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn agent_job_summary_for_user(
|
|
||||||
&self,
|
|
||||||
user_id: &str,
|
|
||||||
) -> Result<AgentJobSummary, DatabaseError> {
|
|
||||||
let conn = self.connect().await?;
|
|
||||||
let mut rows = conn
|
|
||||||
.query(
|
|
||||||
"SELECT status, COUNT(*) as cnt FROM agent_jobs WHERE source = 'direct' AND user_id = ?1 GROUP BY status",
|
|
||||||
params![user_id],
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
|
||||||
|
|
||||||
let mut summary = AgentJobSummary::default();
|
|
||||||
while let Some(row) = rows
|
|
||||||
.next()
|
|
||||||
.await
|
|
||||||
.map_err(|e| DatabaseError::Query(e.to_string()))?
|
|
||||||
{
|
|
||||||
let status = get_text(&row, 0);
|
|
||||||
let count = get_i64(&row, 1) as usize;
|
|
||||||
summary.add_count(&status, count);
|
|
||||||
}
|
|
||||||
Ok(summary)
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn save_action(&self, job_id: Uuid, action: &ActionRecord) -> Result<(), DatabaseError> {
|
async fn save_action(&self, job_id: Uuid, action: &ActionRecord) -> Result<(), DatabaseError> {
|
||||||
let conn = self.connect().await?;
|
let conn = self.connect().await?;
|
||||||
let duration_ms = action.duration.as_millis() as i64;
|
let duration_ms = action.duration.as_millis() as i64;
|
||||||
|
|||||||
@@ -409,15 +409,7 @@ pub trait JobStore: Send + Sync {
|
|||||||
async fn mark_job_stuck(&self, id: Uuid) -> Result<(), DatabaseError>;
|
async fn mark_job_stuck(&self, id: Uuid) -> Result<(), DatabaseError>;
|
||||||
async fn get_stuck_jobs(&self) -> Result<Vec<Uuid>, DatabaseError>;
|
async fn get_stuck_jobs(&self) -> Result<Vec<Uuid>, DatabaseError>;
|
||||||
async fn list_agent_jobs(&self) -> Result<Vec<AgentJobRecord>, DatabaseError>;
|
async fn list_agent_jobs(&self) -> Result<Vec<AgentJobRecord>, DatabaseError>;
|
||||||
async fn list_agent_jobs_for_user(
|
|
||||||
&self,
|
|
||||||
user_id: &str,
|
|
||||||
) -> Result<Vec<AgentJobRecord>, DatabaseError>;
|
|
||||||
async fn agent_job_summary(&self) -> Result<AgentJobSummary, DatabaseError>;
|
async fn agent_job_summary(&self) -> Result<AgentJobSummary, DatabaseError>;
|
||||||
async fn agent_job_summary_for_user(
|
|
||||||
&self,
|
|
||||||
user_id: &str,
|
|
||||||
) -> Result<AgentJobSummary, DatabaseError>;
|
|
||||||
/// Get the failure reason for a single agent job (O(1) lookup).
|
/// Get the failure reason for a single agent job (O(1) lookup).
|
||||||
async fn get_agent_job_failure_reason(&self, id: Uuid)
|
async fn get_agent_job_failure_reason(&self, id: Uuid)
|
||||||
-> Result<Option<String>, DatabaseError>;
|
-> Result<Option<String>, DatabaseError>;
|
||||||
|
|||||||
@@ -249,24 +249,10 @@ impl JobStore for PgBackend {
|
|||||||
self.store.list_agent_jobs().await
|
self.store.list_agent_jobs().await
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn list_agent_jobs_for_user(
|
|
||||||
&self,
|
|
||||||
user_id: &str,
|
|
||||||
) -> Result<Vec<AgentJobRecord>, DatabaseError> {
|
|
||||||
self.store.list_agent_jobs_for_user(user_id).await
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn agent_job_summary(&self) -> Result<AgentJobSummary, DatabaseError> {
|
async fn agent_job_summary(&self) -> Result<AgentJobSummary, DatabaseError> {
|
||||||
self.store.agent_job_summary().await
|
self.store.agent_job_summary().await
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn agent_job_summary_for_user(
|
|
||||||
&self,
|
|
||||||
user_id: &str,
|
|
||||||
) -> Result<AgentJobSummary, DatabaseError> {
|
|
||||||
self.store.agent_job_summary_for_user(user_id).await
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn get_agent_job_failure_reason(
|
async fn get_agent_job_failure_reason(
|
||||||
&self,
|
&self,
|
||||||
id: Uuid,
|
id: Uuid,
|
||||||
|
|||||||
+233
-326
File diff suppressed because it is too large
Load Diff
@@ -842,38 +842,6 @@ impl Store {
|
|||||||
.collect())
|
.collect())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn list_agent_jobs_for_user(
|
|
||||||
&self,
|
|
||||||
user_id: &str,
|
|
||||||
) -> Result<Vec<AgentJobRecord>, DatabaseError> {
|
|
||||||
let conn = self.conn().await?;
|
|
||||||
let rows = conn
|
|
||||||
.query(
|
|
||||||
r#"
|
|
||||||
SELECT id, title, status, user_id, failure_reason,
|
|
||||||
created_at, started_at, completed_at
|
|
||||||
FROM agent_jobs WHERE source = 'direct' AND user_id = $1
|
|
||||||
ORDER BY created_at DESC
|
|
||||||
"#,
|
|
||||||
&[&user_id],
|
|
||||||
)
|
|
||||||
.await?;
|
|
||||||
|
|
||||||
Ok(rows
|
|
||||||
.iter()
|
|
||||||
.map(|r| AgentJobRecord {
|
|
||||||
id: r.get("id"),
|
|
||||||
title: r.get("title"),
|
|
||||||
status: r.get("status"),
|
|
||||||
user_id: r.get::<_, Option<String>>("user_id").unwrap_or_default(),
|
|
||||||
created_at: r.get("created_at"),
|
|
||||||
started_at: r.get("started_at"),
|
|
||||||
completed_at: r.get("completed_at"),
|
|
||||||
failure_reason: r.get("failure_reason"),
|
|
||||||
})
|
|
||||||
.collect())
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Get the failure reason for a single agent job.
|
/// Get the failure reason for a single agent job.
|
||||||
pub async fn get_agent_job_failure_reason(
|
pub async fn get_agent_job_failure_reason(
|
||||||
&self,
|
&self,
|
||||||
@@ -907,27 +875,6 @@ impl Store {
|
|||||||
}
|
}
|
||||||
Ok(summary)
|
Ok(summary)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn agent_job_summary_for_user(
|
|
||||||
&self,
|
|
||||||
user_id: &str,
|
|
||||||
) -> Result<AgentJobSummary, DatabaseError> {
|
|
||||||
let conn = self.conn().await?;
|
|
||||||
let rows = conn
|
|
||||||
.query(
|
|
||||||
"SELECT status, COUNT(*) as cnt FROM agent_jobs WHERE source = 'direct' AND user_id = $1 GROUP BY status",
|
|
||||||
&[&user_id],
|
|
||||||
)
|
|
||||||
.await?;
|
|
||||||
|
|
||||||
let mut summary = AgentJobSummary::default();
|
|
||||||
for row in &rows {
|
|
||||||
let status: String = row.get("status");
|
|
||||||
let count: i64 = row.get("cnt");
|
|
||||||
summary.add_count(&status, count as usize);
|
|
||||||
}
|
|
||||||
Ok(summary)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// ==================== Job Events ====================
|
// ==================== Job Events ====================
|
||||||
|
|||||||
+15
-72
@@ -142,11 +142,6 @@ async fn async_main() -> anyhow::Result<()> {
|
|||||||
init_cli_tracing();
|
init_cli_tracing();
|
||||||
return ironclaw::cli::run_logs_command(logs_cmd.clone(), cli.config.as_deref()).await;
|
return ironclaw::cli::run_logs_command(logs_cmd.clone(), cli.config.as_deref()).await;
|
||||||
}
|
}
|
||||||
Some(Command::Models(models_cmd)) => {
|
|
||||||
init_cli_tracing();
|
|
||||||
return ironclaw::cli::run_models_command(models_cmd.clone(), cli.config.as_deref())
|
|
||||||
.await;
|
|
||||||
}
|
|
||||||
Some(Command::Doctor) => {
|
Some(Command::Doctor) => {
|
||||||
init_cli_tracing();
|
init_cli_tracing();
|
||||||
return ironclaw::cli::run_doctor_command().await;
|
return ironclaw::cli::run_doctor_command().await;
|
||||||
@@ -589,48 +584,15 @@ async fn async_main() -> anyhow::Result<()> {
|
|||||||
// ── Gateway channel ────────────────────────────────────────────────
|
// ── Gateway channel ────────────────────────────────────────────────
|
||||||
|
|
||||||
let mut gateway_url: Option<String> = None;
|
let mut gateway_url: Option<String> = None;
|
||||||
let mut sse_manager: Option<std::sync::Arc<ironclaw::channels::web::sse::SseManager>> = None;
|
let mut sse_sender: Option<
|
||||||
let mut _gateway_state: Option<std::sync::Arc<ironclaw::channels::web::server::GatewayState>> =
|
tokio::sync::broadcast::Sender<ironclaw::channels::web::types::SseEvent>,
|
||||||
None;
|
> = None;
|
||||||
if let Some(ref gw_config) = config.channels.gateway {
|
if let Some(ref gw_config) = config.channels.gateway {
|
||||||
// Build multi-user auth state if user_tokens is configured, else single-user.
|
let mut gw =
|
||||||
let mut gw = if let Some(ref user_tokens) = gw_config.user_tokens {
|
GatewayChannel::new(gw_config.clone()).with_llm_provider(Arc::clone(&components.llm));
|
||||||
use ironclaw::channels::web::auth::{MultiAuthState, UserIdentity};
|
|
||||||
let tokens = user_tokens
|
|
||||||
.iter()
|
|
||||||
.map(|(token, cfg)| {
|
|
||||||
(
|
|
||||||
token.clone(),
|
|
||||||
UserIdentity {
|
|
||||||
user_id: cfg.user_id.clone(),
|
|
||||||
workspace_read_scopes: cfg.workspace_read_scopes.clone(),
|
|
||||||
},
|
|
||||||
)
|
|
||||||
})
|
|
||||||
.collect();
|
|
||||||
let auth = MultiAuthState::multi(tokens);
|
|
||||||
GatewayChannel::new_multi_auth(gw_config.clone(), auth)
|
|
||||||
} else {
|
|
||||||
GatewayChannel::new(gw_config.clone())
|
|
||||||
};
|
|
||||||
gw = gw.with_llm_provider(Arc::clone(&components.llm));
|
|
||||||
if let Some(ref ws) = components.workspace {
|
if let Some(ref ws) = components.workspace {
|
||||||
gw = gw.with_workspace(Arc::clone(ws));
|
gw = gw.with_workspace(Arc::clone(ws));
|
||||||
}
|
}
|
||||||
// Create per-user workspace pool for multi-user mode.
|
|
||||||
if let Some(ref db) = components.db {
|
|
||||||
let emb_cache_config = ironclaw::workspace::EmbeddingCacheConfig {
|
|
||||||
max_entries: config.embeddings.cache_size,
|
|
||||||
};
|
|
||||||
let pool = Arc::new(ironclaw::channels::web::server::WorkspacePool::new(
|
|
||||||
Arc::clone(db),
|
|
||||||
components.embeddings.clone(),
|
|
||||||
emb_cache_config,
|
|
||||||
config.search.clone(),
|
|
||||||
config.workspace.clone(),
|
|
||||||
));
|
|
||||||
gw = gw.with_workspace_pool(pool);
|
|
||||||
}
|
|
||||||
gw = gw.with_session_manager(Arc::clone(&session_manager));
|
gw = gw.with_session_manager(Arc::clone(&session_manager));
|
||||||
gw = gw.with_log_broadcaster(Arc::clone(&log_broadcaster));
|
gw = gw.with_log_broadcaster(Arc::clone(&log_broadcaster));
|
||||||
gw = gw.with_log_level_handle(Arc::clone(&log_level_handle));
|
gw = gw.with_log_level_handle(Arc::clone(&log_level_handle));
|
||||||
@@ -681,12 +643,8 @@ async fn async_main() -> anyhow::Result<()> {
|
|||||||
let mut rx = tx.subscribe();
|
let mut rx = tx.subscribe();
|
||||||
let gw_state = Arc::clone(gw.state());
|
let gw_state = Arc::clone(gw.state());
|
||||||
tokio::spawn(async move {
|
tokio::spawn(async move {
|
||||||
while let Ok((_job_id, user_id, event)) = rx.recv().await {
|
while let Ok((_job_id, event)) = rx.recv().await {
|
||||||
if user_id.is_empty() {
|
gw_state.sse.broadcast(event);
|
||||||
gw_state.sse.broadcast(event);
|
|
||||||
} else {
|
|
||||||
gw_state.sse.broadcast_for_user(&user_id, event);
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
@@ -728,8 +686,7 @@ async fn async_main() -> anyhow::Result<()> {
|
|||||||
// Capture SSE sender and routine engine slot before moving gw into channels.
|
// Capture SSE sender and routine engine slot before moving gw into channels.
|
||||||
// IMPORTANT: This must come after all `with_*` calls since `rebuild_state`
|
// IMPORTANT: This must come after all `with_*` calls since `rebuild_state`
|
||||||
// creates a new SseManager, which would orphan this sender.
|
// creates a new SseManager, which would orphan this sender.
|
||||||
sse_manager = Some(Arc::clone(&gw.state().sse));
|
sse_sender = Some(gw.state().sse.sender());
|
||||||
_gateway_state = Some(Arc::clone(gw.state()));
|
|
||||||
channel_names.push("gateway".to_string());
|
channel_names.push("gateway".to_string());
|
||||||
channels.add(Box::new(gw)).await;
|
channels.add(Box::new(gw)).await;
|
||||||
}
|
}
|
||||||
@@ -812,20 +769,12 @@ async fn async_main() -> anyhow::Result<()> {
|
|||||||
|
|
||||||
// Auto-activate WASM channels that were active in a previous session.
|
// Auto-activate WASM channels that were active in a previous session.
|
||||||
// Relay channels are handled separately below via restore_relay_channels().
|
// Relay channels are handled separately below via restore_relay_channels().
|
||||||
let ext_user_id = config
|
let persisted = ext_mgr.load_persisted_active_channels().await;
|
||||||
.channels
|
|
||||||
.gateway
|
|
||||||
.as_ref()
|
|
||||||
.map(|g| g.user_id.clone())
|
|
||||||
.unwrap_or_else(|| "default".to_string());
|
|
||||||
let persisted = ext_mgr.load_persisted_active_channels(&ext_user_id).await;
|
|
||||||
for name in &persisted {
|
for name in &persisted {
|
||||||
if active_at_startup.contains(name)
|
if active_at_startup.contains(name) || ext_mgr.is_relay_channel(name).await {
|
||||||
|| ext_mgr.is_relay_channel(name, &ext_user_id).await
|
|
||||||
{
|
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
match ext_mgr.activate(name, &ext_user_id).await {
|
match ext_mgr.activate(name).await {
|
||||||
Ok(result) => {
|
Ok(result) => {
|
||||||
tracing::debug!(
|
tracing::debug!(
|
||||||
channel = %name,
|
channel = %name,
|
||||||
@@ -850,20 +799,14 @@ async fn async_main() -> anyhow::Result<()> {
|
|||||||
ext_mgr
|
ext_mgr
|
||||||
.set_relay_channel_manager(Arc::clone(&channels))
|
.set_relay_channel_manager(Arc::clone(&channels))
|
||||||
.await;
|
.await;
|
||||||
let ext_user_id = config
|
ext_mgr.restore_relay_channels().await;
|
||||||
.channels
|
|
||||||
.gateway
|
|
||||||
.as_ref()
|
|
||||||
.map(|g| g.user_id.clone())
|
|
||||||
.unwrap_or_else(|| "default".to_string());
|
|
||||||
ext_mgr.restore_relay_channels(&ext_user_id).await;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Wire SSE sender into extension manager for broadcasting status events.
|
// Wire SSE sender into extension manager for broadcasting status events.
|
||||||
if let Some(ref ext_mgr) = components.extension_manager
|
if let Some(ref ext_mgr) = components.extension_manager
|
||||||
&& let Some(sse) = sse_manager
|
&& let Some(ref sender) = sse_sender
|
||||||
{
|
{
|
||||||
ext_mgr.set_sse_sender(sse).await;
|
ext_mgr.set_sse_sender(sender.clone()).await;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Snapshot memory for trace recording before the agent starts
|
// Snapshot memory for trace recording before the agent starts
|
||||||
@@ -901,7 +844,7 @@ async fn async_main() -> anyhow::Result<()> {
|
|||||||
skills_config: config.skills.clone(),
|
skills_config: config.skills.clone(),
|
||||||
hooks: components.hooks,
|
hooks: components.hooks,
|
||||||
cost_guard: components.cost_guard,
|
cost_guard: components.cost_guard,
|
||||||
sse_tx: None, // TODO: wire SseManager into scheduler (needs Sender<SseEvent> → Arc<SseManager> refactor)
|
sse_tx: sse_sender,
|
||||||
http_interceptor,
|
http_interceptor,
|
||||||
transcription: config.transcription.create_provider().map(|p| {
|
transcription: config.transcription.create_provider().map(|p| {
|
||||||
Arc::new(ironclaw::llm::transcription::TranscriptionMiddleware::new(
|
Arc::new(ironclaw::llm::transcription::TranscriptionMiddleware::new(
|
||||||
|
|||||||
+6
-24
@@ -40,8 +40,7 @@ pub struct OrchestratorState {
|
|||||||
pub job_manager: Arc<ContainerJobManager>,
|
pub job_manager: Arc<ContainerJobManager>,
|
||||||
pub token_store: TokenStore,
|
pub token_store: TokenStore,
|
||||||
/// Broadcast channel for job events (consumed by the web gateway SSE).
|
/// Broadcast channel for job events (consumed by the web gateway SSE).
|
||||||
/// Tuple: (job_id, user_id, event).
|
pub job_event_tx: Option<broadcast::Sender<(Uuid, SseEvent)>>,
|
||||||
pub job_event_tx: Option<broadcast::Sender<(Uuid, String, SseEvent)>>,
|
|
||||||
/// Buffered follow-up prompts for sandbox jobs, keyed by job_id.
|
/// Buffered follow-up prompts for sandbox jobs, keyed by job_id.
|
||||||
pub prompt_queue: Arc<Mutex<HashMap<Uuid, VecDeque<PendingPrompt>>>>,
|
pub prompt_queue: Arc<Mutex<HashMap<Uuid, VecDeque<PendingPrompt>>>>,
|
||||||
/// Database handle for persisting job events.
|
/// Database handle for persisting job events.
|
||||||
@@ -352,24 +351,9 @@ async fn job_event_handler(
|
|||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
// Broadcast via the channel (if configured).
|
// Broadcast via the channel (if configured)
|
||||||
// Look up the job owner so the gateway can scope delivery per-user.
|
|
||||||
if let Some(ref tx) = state.job_event_tx {
|
if let Some(ref tx) = state.job_event_tx {
|
||||||
let user_id = match state.store.as_ref() {
|
let _ = tx.send((job_id, sse_event));
|
||||||
Some(store) => store
|
|
||||||
.get_sandbox_job(job_id)
|
|
||||||
.await
|
|
||||||
.ok()
|
|
||||||
.flatten()
|
|
||||||
.map(|j| j.user_id),
|
|
||||||
None => None,
|
|
||||||
};
|
|
||||||
if let Some(uid) = user_id {
|
|
||||||
let _ = tx.send((job_id, uid, sse_event));
|
|
||||||
} else {
|
|
||||||
// Fallback: broadcast globally (single-user mode or job not found).
|
|
||||||
let _ = tx.send((job_id, String::new(), sse_event));
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(StatusCode::OK)
|
Ok(StatusCode::OK)
|
||||||
@@ -785,10 +769,8 @@ mod tests {
|
|||||||
let resp = router.oneshot(req).await.unwrap();
|
let resp = router.oneshot(req).await.unwrap();
|
||||||
assert_eq!(resp.status(), StatusCode::OK);
|
assert_eq!(resp.status(), StatusCode::OK);
|
||||||
|
|
||||||
let (recv_id, recv_uid, event) = rx.recv().await.unwrap();
|
let (recv_id, event) = rx.recv().await.unwrap();
|
||||||
assert_eq!(recv_id, job_id);
|
assert_eq!(recv_id, job_id);
|
||||||
// No store configured, so user_id falls back to empty string.
|
|
||||||
assert_eq!(recv_uid, "");
|
|
||||||
match event {
|
match event {
|
||||||
SseEvent::JobMessage {
|
SseEvent::JobMessage {
|
||||||
job_id: jid,
|
job_id: jid,
|
||||||
@@ -842,7 +824,7 @@ mod tests {
|
|||||||
let resp = router.oneshot(req).await.unwrap();
|
let resp = router.oneshot(req).await.unwrap();
|
||||||
assert_eq!(resp.status(), StatusCode::OK);
|
assert_eq!(resp.status(), StatusCode::OK);
|
||||||
|
|
||||||
let (_recv_id, _recv_uid, event) = rx.recv().await.unwrap();
|
let (_recv_id, event) = rx.recv().await.unwrap();
|
||||||
match event {
|
match event {
|
||||||
SseEvent::JobToolUse { tool_name, .. } => {
|
SseEvent::JobToolUse { tool_name, .. } => {
|
||||||
assert_eq!(tool_name, "shell");
|
assert_eq!(tool_name, "shell");
|
||||||
@@ -887,7 +869,7 @@ mod tests {
|
|||||||
let resp = router.oneshot(req).await.unwrap();
|
let resp = router.oneshot(req).await.unwrap();
|
||||||
assert_eq!(resp.status(), StatusCode::OK);
|
assert_eq!(resp.status(), StatusCode::OK);
|
||||||
|
|
||||||
let (_recv_id, _recv_uid, event) = rx.recv().await.unwrap();
|
let (_recv_id, event) = rx.recv().await.unwrap();
|
||||||
// Unknown event types fall through to JobStatus
|
// Unknown event types fall through to JobStatus
|
||||||
assert!(matches!(event, SseEvent::JobStatus { .. }));
|
assert!(matches!(event, SseEvent::JobStatus { .. }));
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -63,7 +63,7 @@ fn resolve_orchestrator_port() -> u16 {
|
|||||||
/// Result of orchestrator setup, containing all handles needed by the agent.
|
/// Result of orchestrator setup, containing all handles needed by the agent.
|
||||||
pub struct OrchestratorSetup {
|
pub struct OrchestratorSetup {
|
||||||
pub container_job_manager: Option<Arc<ContainerJobManager>>,
|
pub container_job_manager: Option<Arc<ContainerJobManager>>,
|
||||||
pub job_event_tx: Option<broadcast::Sender<(Uuid, String, SseEvent)>>,
|
pub job_event_tx: Option<broadcast::Sender<(Uuid, SseEvent)>>,
|
||||||
pub prompt_queue: Arc<Mutex<HashMap<Uuid, VecDeque<api::PendingPrompt>>>>,
|
pub prompt_queue: Arc<Mutex<HashMap<Uuid, VecDeque<api::PendingPrompt>>>>,
|
||||||
pub docker_status: crate::sandbox::DockerStatus,
|
pub docker_status: crate::sandbox::DockerStatus,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -130,7 +130,7 @@ impl Tool for ToolInstallTool {
|
|||||||
async fn execute(
|
async fn execute(
|
||||||
&self,
|
&self,
|
||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
ctx: &JobContext,
|
_ctx: &JobContext,
|
||||||
) -> Result<ToolOutput, ToolError> {
|
) -> Result<ToolOutput, ToolError> {
|
||||||
let start = std::time::Instant::now();
|
let start = std::time::Instant::now();
|
||||||
|
|
||||||
@@ -150,7 +150,7 @@ impl Tool for ToolInstallTool {
|
|||||||
|
|
||||||
let result = self
|
let result = self
|
||||||
.manager
|
.manager
|
||||||
.install(name, url, kind_hint, &ctx.user_id)
|
.install(name, url, kind_hint)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
||||||
|
|
||||||
@@ -205,7 +205,7 @@ impl Tool for ToolAuthTool {
|
|||||||
async fn execute(
|
async fn execute(
|
||||||
&self,
|
&self,
|
||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
ctx: &JobContext,
|
_ctx: &JobContext,
|
||||||
) -> Result<ToolOutput, ToolError> {
|
) -> Result<ToolOutput, ToolError> {
|
||||||
let start = std::time::Instant::now();
|
let start = std::time::Instant::now();
|
||||||
|
|
||||||
@@ -213,13 +213,13 @@ impl Tool for ToolAuthTool {
|
|||||||
|
|
||||||
let result = self
|
let result = self
|
||||||
.manager
|
.manager
|
||||||
.auth(name, &ctx.user_id)
|
.auth(name)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
||||||
|
|
||||||
// Auto-activate after successful auth so tools are available immediately
|
// Auto-activate after successful auth so tools are available immediately
|
||||||
if result.is_authenticated() {
|
if result.is_authenticated() {
|
||||||
match self.manager.activate(name, &ctx.user_id).await {
|
match self.manager.activate(name).await {
|
||||||
Ok(activate_result) => {
|
Ok(activate_result) => {
|
||||||
let output = serde_json::json!({
|
let output = serde_json::json!({
|
||||||
"status": "authenticated_and_activated",
|
"status": "authenticated_and_activated",
|
||||||
@@ -304,13 +304,13 @@ impl Tool for ToolActivateTool {
|
|||||||
async fn execute(
|
async fn execute(
|
||||||
&self,
|
&self,
|
||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
ctx: &JobContext,
|
_ctx: &JobContext,
|
||||||
) -> Result<ToolOutput, ToolError> {
|
) -> Result<ToolOutput, ToolError> {
|
||||||
let start = std::time::Instant::now();
|
let start = std::time::Instant::now();
|
||||||
|
|
||||||
let name = require_str(¶ms, "name")?;
|
let name = require_str(¶ms, "name")?;
|
||||||
|
|
||||||
match self.manager.activate(name, &ctx.user_id).await {
|
match self.manager.activate(name).await {
|
||||||
Ok(result) => {
|
Ok(result) => {
|
||||||
let output = serde_json::to_value(&result)
|
let output = serde_json::to_value(&result)
|
||||||
.unwrap_or_else(|_| serde_json::json!({"error": "serialization failed"}));
|
.unwrap_or_else(|_| serde_json::json!({"error": "serialization failed"}));
|
||||||
@@ -329,12 +329,12 @@ impl Tool for ToolActivateTool {
|
|||||||
|
|
||||||
// Activation failed due to missing auth; initiate auth flow
|
// Activation failed due to missing auth; initiate auth flow
|
||||||
// so the agent loop can show the auth card.
|
// so the agent loop can show the auth card.
|
||||||
match self.manager.auth(name, &ctx.user_id).await {
|
match self.manager.auth(name).await {
|
||||||
Ok(auth_result) if auth_result.is_authenticated() => {
|
Ok(auth_result) if auth_result.is_authenticated() => {
|
||||||
// Auth succeeded (e.g. env var was set); retry activation.
|
// Auth succeeded (e.g. env var was set); retry activation.
|
||||||
let result = self
|
let result = self
|
||||||
.manager
|
.manager
|
||||||
.activate(name, &ctx.user_id)
|
.activate(name)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
||||||
let output = serde_json::to_value(&result).unwrap_or_else(
|
let output = serde_json::to_value(&result).unwrap_or_else(
|
||||||
@@ -404,7 +404,7 @@ impl Tool for ToolListTool {
|
|||||||
async fn execute(
|
async fn execute(
|
||||||
&self,
|
&self,
|
||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
ctx: &JobContext,
|
_ctx: &JobContext,
|
||||||
) -> Result<ToolOutput, ToolError> {
|
) -> Result<ToolOutput, ToolError> {
|
||||||
let start = std::time::Instant::now();
|
let start = std::time::Instant::now();
|
||||||
|
|
||||||
@@ -425,7 +425,7 @@ impl Tool for ToolListTool {
|
|||||||
|
|
||||||
let extensions = self
|
let extensions = self
|
||||||
.manager
|
.manager
|
||||||
.list(kind_filter, include_available, &ctx.user_id)
|
.list(kind_filter, include_available)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
||||||
|
|
||||||
@@ -477,7 +477,7 @@ impl Tool for ToolRemoveTool {
|
|||||||
async fn execute(
|
async fn execute(
|
||||||
&self,
|
&self,
|
||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
ctx: &JobContext,
|
_ctx: &JobContext,
|
||||||
) -> Result<ToolOutput, ToolError> {
|
) -> Result<ToolOutput, ToolError> {
|
||||||
let start = std::time::Instant::now();
|
let start = std::time::Instant::now();
|
||||||
|
|
||||||
@@ -485,7 +485,7 @@ impl Tool for ToolRemoveTool {
|
|||||||
|
|
||||||
let message = self
|
let message = self
|
||||||
.manager
|
.manager
|
||||||
.remove(name, &ctx.user_id)
|
.remove(name)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
||||||
|
|
||||||
@@ -541,7 +541,7 @@ impl Tool for ToolUpgradeTool {
|
|||||||
async fn execute(
|
async fn execute(
|
||||||
&self,
|
&self,
|
||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
ctx: &JobContext,
|
_ctx: &JobContext,
|
||||||
) -> Result<ToolOutput, ToolError> {
|
) -> Result<ToolOutput, ToolError> {
|
||||||
let start = std::time::Instant::now();
|
let start = std::time::Instant::now();
|
||||||
|
|
||||||
@@ -549,7 +549,7 @@ impl Tool for ToolUpgradeTool {
|
|||||||
|
|
||||||
let result = self
|
let result = self
|
||||||
.manager
|
.manager
|
||||||
.upgrade(name, &ctx.user_id)
|
.upgrade(name)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
||||||
|
|
||||||
@@ -603,7 +603,7 @@ impl Tool for ExtensionInfoTool {
|
|||||||
async fn execute(
|
async fn execute(
|
||||||
&self,
|
&self,
|
||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
ctx: &JobContext,
|
_ctx: &JobContext,
|
||||||
) -> Result<ToolOutput, ToolError> {
|
) -> Result<ToolOutput, ToolError> {
|
||||||
let start = std::time::Instant::now();
|
let start = std::time::Instant::now();
|
||||||
|
|
||||||
@@ -611,7 +611,7 @@ impl Tool for ExtensionInfoTool {
|
|||||||
|
|
||||||
let info = self
|
let info = self
|
||||||
.manager
|
.manager
|
||||||
.extension_info(name, &ctx.user_id)
|
.extension_info(name)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
||||||
|
|
||||||
|
|||||||
@@ -85,7 +85,7 @@ pub struct CreateJobTool {
|
|||||||
job_manager: Option<Arc<ContainerJobManager>>,
|
job_manager: Option<Arc<ContainerJobManager>>,
|
||||||
store: Option<Arc<dyn Database>>,
|
store: Option<Arc<dyn Database>>,
|
||||||
/// Broadcast sender for job events (used to subscribe a monitor).
|
/// Broadcast sender for job events (used to subscribe a monitor).
|
||||||
event_tx: Option<tokio::sync::broadcast::Sender<(Uuid, String, SseEvent)>>,
|
event_tx: Option<tokio::sync::broadcast::Sender<(Uuid, SseEvent)>>,
|
||||||
/// Injection channel for pushing messages into the agent loop.
|
/// Injection channel for pushing messages into the agent loop.
|
||||||
inject_tx: Option<tokio::sync::mpsc::Sender<IncomingMessage>>,
|
inject_tx: Option<tokio::sync::mpsc::Sender<IncomingMessage>>,
|
||||||
/// Encrypted secrets store for validating credential grants.
|
/// Encrypted secrets store for validating credential grants.
|
||||||
@@ -120,7 +120,7 @@ impl CreateJobTool {
|
|||||||
/// monitor that forwards Claude Code output to the main agent loop.
|
/// monitor that forwards Claude Code output to the main agent loop.
|
||||||
pub fn with_monitor_deps(
|
pub fn with_monitor_deps(
|
||||||
mut self,
|
mut self,
|
||||||
event_tx: tokio::sync::broadcast::Sender<(Uuid, String, SseEvent)>,
|
event_tx: tokio::sync::broadcast::Sender<(Uuid, SseEvent)>,
|
||||||
inject_tx: tokio::sync::mpsc::Sender<IncomingMessage>,
|
inject_tx: tokio::sync::mpsc::Sender<IncomingMessage>,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
self.event_tx = Some(event_tx);
|
self.event_tx = Some(event_tx);
|
||||||
|
|||||||
+45
-358
@@ -12,119 +12,15 @@
|
|||||||
//! Use `memory_write` to persist important facts that should be remembered
|
//! Use `memory_write` to persist important facts that should be remembered
|
||||||
//! across sessions.
|
//! across sessions.
|
||||||
|
|
||||||
use std::collections::HashMap;
|
|
||||||
use std::path::Path;
|
use std::path::Path;
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
|
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
use tokio::sync::RwLock;
|
|
||||||
|
|
||||||
use crate::context::JobContext;
|
use crate::context::JobContext;
|
||||||
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
|
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
|
||||||
use crate::workspace::{Workspace, paths};
|
use crate::workspace::{Workspace, paths};
|
||||||
|
|
||||||
// ── WorkspaceResolver ──────────────────────────────────────────────
|
|
||||||
|
|
||||||
/// Resolves a workspace for a given user ID.
|
|
||||||
///
|
|
||||||
/// In single-user mode, always returns the same workspace.
|
|
||||||
/// In multi-tenant mode, creates per-user workspaces on demand.
|
|
||||||
#[async_trait]
|
|
||||||
pub trait WorkspaceResolver: Send + Sync {
|
|
||||||
async fn resolve(&self, user_id: &str) -> Arc<Workspace>;
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Returns a fixed workspace regardless of user ID (single-user mode).
|
|
||||||
pub struct FixedWorkspaceResolver {
|
|
||||||
workspace: Arc<Workspace>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl FixedWorkspaceResolver {
|
|
||||||
pub fn new(workspace: Arc<Workspace>) -> Self {
|
|
||||||
Self { workspace }
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[async_trait]
|
|
||||||
impl WorkspaceResolver for FixedWorkspaceResolver {
|
|
||||||
async fn resolve(&self, _user_id: &str) -> Arc<Workspace> {
|
|
||||||
Arc::clone(&self.workspace)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Creates per-user workspaces on demand, caching them for reuse.
|
|
||||||
///
|
|
||||||
/// Used in multi-tenant mode where each authenticated user gets their own
|
|
||||||
/// workspace scope. The workspace is constructed with the same configuration
|
|
||||||
/// (embeddings, search config, memory layers) as the startup workspace.
|
|
||||||
pub struct PerUserWorkspaceResolver {
|
|
||||||
db: Arc<dyn crate::db::Database>,
|
|
||||||
embeddings: Option<Arc<dyn crate::workspace::EmbeddingProvider>>,
|
|
||||||
embedding_cache_config: crate::workspace::EmbeddingCacheConfig,
|
|
||||||
search_config: crate::config::WorkspaceSearchConfig,
|
|
||||||
workspace_config: crate::config::WorkspaceConfig,
|
|
||||||
cache: RwLock<HashMap<String, Arc<Workspace>>>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl PerUserWorkspaceResolver {
|
|
||||||
pub fn new(
|
|
||||||
db: Arc<dyn crate::db::Database>,
|
|
||||||
embeddings: Option<Arc<dyn crate::workspace::EmbeddingProvider>>,
|
|
||||||
embedding_cache_config: crate::workspace::EmbeddingCacheConfig,
|
|
||||||
search_config: crate::config::WorkspaceSearchConfig,
|
|
||||||
workspace_config: crate::config::WorkspaceConfig,
|
|
||||||
) -> Self {
|
|
||||||
Self {
|
|
||||||
db,
|
|
||||||
embeddings,
|
|
||||||
embedding_cache_config,
|
|
||||||
search_config,
|
|
||||||
workspace_config,
|
|
||||||
cache: RwLock::new(HashMap::new()),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn build_workspace(&self, user_id: &str) -> Arc<Workspace> {
|
|
||||||
let mut ws = Workspace::new_with_db(user_id, Arc::clone(&self.db))
|
|
||||||
.with_search_config(&self.search_config);
|
|
||||||
|
|
||||||
if let Some(ref emb) = self.embeddings {
|
|
||||||
ws = ws.with_embeddings_cached(Arc::clone(emb), self.embedding_cache_config.clone());
|
|
||||||
}
|
|
||||||
|
|
||||||
if !self.workspace_config.read_scopes.is_empty() {
|
|
||||||
ws = ws.with_additional_read_scopes(self.workspace_config.read_scopes.clone());
|
|
||||||
}
|
|
||||||
ws = ws.with_memory_layers(self.workspace_config.memory_layers.clone());
|
|
||||||
|
|
||||||
Arc::new(ws)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[async_trait]
|
|
||||||
impl WorkspaceResolver for PerUserWorkspaceResolver {
|
|
||||||
async fn resolve(&self, user_id: &str) -> Arc<Workspace> {
|
|
||||||
// Fast path: read lock
|
|
||||||
{
|
|
||||||
let cache = self.cache.read().await;
|
|
||||||
if let Some(ws) = cache.get(user_id) {
|
|
||||||
return Arc::clone(ws);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Slow path: write lock, double-check
|
|
||||||
let mut cache = self.cache.write().await;
|
|
||||||
if let Some(ws) = cache.get(user_id) {
|
|
||||||
return Arc::clone(ws);
|
|
||||||
}
|
|
||||||
|
|
||||||
let ws = self.build_workspace(user_id);
|
|
||||||
cache.insert(user_id.to_string(), Arc::clone(&ws));
|
|
||||||
tracing::debug!(user_id = user_id, "Created per-user workspace");
|
|
||||||
ws
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Detect paths that are clearly local filesystem references, not workspace-memory docs.
|
/// Detect paths that are clearly local filesystem references, not workspace-memory docs.
|
||||||
///
|
///
|
||||||
/// Examples:
|
/// Examples:
|
||||||
@@ -166,20 +62,13 @@ fn map_write_err(e: crate::error::WorkspaceError) -> ToolError {
|
|||||||
/// The agent should call this tool before answering questions about
|
/// The agent should call this tool before answering questions about
|
||||||
/// prior work, decisions, preferences, or any historical context.
|
/// prior work, decisions, preferences, or any historical context.
|
||||||
pub struct MemorySearchTool {
|
pub struct MemorySearchTool {
|
||||||
resolver: Arc<dyn WorkspaceResolver>,
|
workspace: Arc<Workspace>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl MemorySearchTool {
|
impl MemorySearchTool {
|
||||||
/// Create a new memory search tool with a workspace resolver.
|
/// Create a new memory search tool.
|
||||||
pub fn new(resolver: Arc<dyn WorkspaceResolver>) -> Self {
|
pub fn new(workspace: Arc<Workspace>) -> Self {
|
||||||
Self { resolver }
|
Self { workspace }
|
||||||
}
|
|
||||||
|
|
||||||
/// Create from a fixed workspace (backward compatibility).
|
|
||||||
pub fn from_workspace(workspace: Arc<Workspace>) -> Self {
|
|
||||||
Self {
|
|
||||||
resolver: Arc::new(FixedWorkspaceResolver::new(workspace)),
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -218,7 +107,7 @@ impl Tool for MemorySearchTool {
|
|||||||
async fn execute(
|
async fn execute(
|
||||||
&self,
|
&self,
|
||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
ctx: &JobContext,
|
_ctx: &JobContext,
|
||||||
) -> Result<ToolOutput, ToolError> {
|
) -> Result<ToolOutput, ToolError> {
|
||||||
let start = std::time::Instant::now();
|
let start = std::time::Instant::now();
|
||||||
|
|
||||||
@@ -230,8 +119,8 @@ impl Tool for MemorySearchTool {
|
|||||||
.unwrap_or(5)
|
.unwrap_or(5)
|
||||||
.min(20) as usize;
|
.min(20) as usize;
|
||||||
|
|
||||||
let workspace = self.resolver.resolve(&ctx.user_id).await;
|
let results = self
|
||||||
let results = workspace
|
.workspace
|
||||||
.search(query, limit)
|
.search(query, limit)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("Search failed: {}", e)))?;
|
.map_err(|e| ToolError::ExecutionFailed(format!("Search failed: {}", e)))?;
|
||||||
@@ -262,20 +151,13 @@ impl Tool for MemorySearchTool {
|
|||||||
/// Use this to persist important information that should be remembered
|
/// Use this to persist important information that should be remembered
|
||||||
/// across sessions: decisions, preferences, facts, lessons learned.
|
/// across sessions: decisions, preferences, facts, lessons learned.
|
||||||
pub struct MemoryWriteTool {
|
pub struct MemoryWriteTool {
|
||||||
resolver: Arc<dyn WorkspaceResolver>,
|
workspace: Arc<Workspace>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl MemoryWriteTool {
|
impl MemoryWriteTool {
|
||||||
/// Create a new memory write tool with a workspace resolver.
|
/// Create a new memory write tool.
|
||||||
pub fn new(resolver: Arc<dyn WorkspaceResolver>) -> Self {
|
pub fn new(workspace: Arc<Workspace>) -> Self {
|
||||||
Self { resolver }
|
Self { workspace }
|
||||||
}
|
|
||||||
|
|
||||||
/// Create from a fixed workspace (backward compatibility).
|
|
||||||
pub fn from_workspace(workspace: Arc<Workspace>) -> Self {
|
|
||||||
Self {
|
|
||||||
resolver: Arc::new(FixedWorkspaceResolver::new(workspace)),
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -349,21 +231,19 @@ impl Tool for MemoryWriteTool {
|
|||||||
)));
|
)));
|
||||||
}
|
}
|
||||||
|
|
||||||
let workspace = self.resolver.resolve(&ctx.user_id).await;
|
|
||||||
|
|
||||||
// Bootstrap target: clear BOOTSTRAP.md to mark first-run ritual complete.
|
// Bootstrap target: clear BOOTSTRAP.md to mark first-run ritual complete.
|
||||||
// Handled early because it accepts empty content (unlike other targets).
|
// Handled early because it accepts empty content (unlike other targets).
|
||||||
if target == "bootstrap" {
|
if target == "bootstrap" {
|
||||||
// Write empty content to effectively disable the bootstrap injection.
|
// Write empty content to effectively disable the bootstrap injection.
|
||||||
// system_prompt_for_context() skips empty files.
|
// system_prompt_for_context() skips empty files.
|
||||||
workspace
|
self.workspace
|
||||||
.write(paths::BOOTSTRAP, "")
|
.write(paths::BOOTSTRAP, "")
|
||||||
.await
|
.await
|
||||||
.map_err(map_write_err)?;
|
.map_err(map_write_err)?;
|
||||||
|
|
||||||
// Also set the in-memory flag so BOOTSTRAP.md injection stops
|
// Also set the in-memory flag so BOOTSTRAP.md injection stops
|
||||||
// immediately without waiting for a restart.
|
// immediately without waiting for a restart.
|
||||||
workspace.mark_bootstrap_completed();
|
self.workspace.mark_bootstrap_completed();
|
||||||
|
|
||||||
let output = serde_json::json!({
|
let output = serde_json::json!({
|
||||||
"status": "cleared",
|
"status": "cleared",
|
||||||
@@ -409,12 +289,12 @@ impl Tool for MemoryWriteTool {
|
|||||||
// Otherwise, use default workspace methods (which include injection scanning).
|
// Otherwise, use default workspace methods (which include injection scanning).
|
||||||
let layer_result = if let Some(layer_name) = layer {
|
let layer_result = if let Some(layer_name) = layer {
|
||||||
let result = if append {
|
let result = if append {
|
||||||
workspace
|
self.workspace
|
||||||
.append_to_layer(layer_name, &resolved_path, content, force)
|
.append_to_layer(layer_name, &resolved_path, content, force)
|
||||||
.await
|
.await
|
||||||
.map_err(map_write_err)?
|
.map_err(map_write_err)?
|
||||||
} else {
|
} else {
|
||||||
workspace
|
self.workspace
|
||||||
.write_to_layer(layer_name, &resolved_path, content, force)
|
.write_to_layer(layer_name, &resolved_path, content, force)
|
||||||
.await
|
.await
|
||||||
.map_err(map_write_err)?
|
.map_err(map_write_err)?
|
||||||
@@ -427,33 +307,31 @@ impl Tool for MemoryWriteTool {
|
|||||||
match target {
|
match target {
|
||||||
"memory" => {
|
"memory" => {
|
||||||
if append {
|
if append {
|
||||||
workspace
|
self.workspace
|
||||||
.append_memory(content)
|
.append_memory(content)
|
||||||
.await
|
.await
|
||||||
.map_err(map_write_err)?;
|
.map_err(map_write_err)?;
|
||||||
} else {
|
} else {
|
||||||
workspace
|
self.workspace
|
||||||
.write(paths::MEMORY, content)
|
.write(paths::MEMORY, content)
|
||||||
.await
|
.await
|
||||||
.map_err(map_write_err)?;
|
.map_err(map_write_err)?;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
"daily_log" => {
|
"daily_log" => {
|
||||||
let tz = crate::timezone::parse_timezone(&ctx.user_timezone)
|
self.workspace
|
||||||
.unwrap_or(chrono_tz::Tz::UTC);
|
|
||||||
workspace
|
|
||||||
.append_daily_log_tz(content, tz)
|
.append_daily_log_tz(content, tz)
|
||||||
.await
|
.await
|
||||||
.map_err(map_write_err)?;
|
.map_err(map_write_err)?;
|
||||||
}
|
}
|
||||||
_ => {
|
_ => {
|
||||||
if append {
|
if append {
|
||||||
workspace
|
self.workspace
|
||||||
.append(&resolved_path, content)
|
.append(&resolved_path, content)
|
||||||
.await
|
.await
|
||||||
.map_err(map_write_err)?;
|
.map_err(map_write_err)?;
|
||||||
} else {
|
} else {
|
||||||
workspace
|
self.workspace
|
||||||
.write(&resolved_path, content)
|
.write(&resolved_path, content)
|
||||||
.await
|
.await
|
||||||
.map_err(map_write_err)?;
|
.map_err(map_write_err)?;
|
||||||
@@ -483,12 +361,12 @@ impl Tool for MemoryWriteTool {
|
|||||||
};
|
};
|
||||||
let mut synced_docs: Vec<&str> = Vec::new();
|
let mut synced_docs: Vec<&str> = Vec::new();
|
||||||
if normalized_path == paths::PROFILE {
|
if normalized_path == paths::PROFILE {
|
||||||
match workspace.sync_profile_documents().await {
|
match self.workspace.sync_profile_documents().await {
|
||||||
Ok(true) => {
|
Ok(true) => {
|
||||||
tracing::info!("profile write: synced USER.md + assistant-directives.md");
|
tracing::info!("profile write: synced USER.md + assistant-directives.md");
|
||||||
synced_docs.extend_from_slice(&[paths::USER, paths::ASSISTANT_DIRECTIVES]);
|
synced_docs.extend_from_slice(&[paths::USER, paths::ASSISTANT_DIRECTIVES]);
|
||||||
|
|
||||||
workspace.mark_bootstrap_completed();
|
self.workspace.mark_bootstrap_completed();
|
||||||
let toml_path = crate::settings::Settings::default_toml_path();
|
let toml_path = crate::settings::Settings::default_toml_path();
|
||||||
if let Ok(Some(mut settings)) = crate::settings::Settings::load_toml(&toml_path)
|
if let Ok(Some(mut settings)) = crate::settings::Settings::load_toml(&toml_path)
|
||||||
&& !settings.profile_onboarding_completed
|
&& !settings.profile_onboarding_completed
|
||||||
@@ -538,20 +416,13 @@ impl Tool for MemoryWriteTool {
|
|||||||
///
|
///
|
||||||
/// Use this to read the full content of any file in the workspace.
|
/// Use this to read the full content of any file in the workspace.
|
||||||
pub struct MemoryReadTool {
|
pub struct MemoryReadTool {
|
||||||
resolver: Arc<dyn WorkspaceResolver>,
|
workspace: Arc<Workspace>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl MemoryReadTool {
|
impl MemoryReadTool {
|
||||||
/// Create a new memory read tool with a workspace resolver.
|
/// Create a new memory read tool.
|
||||||
pub fn new(resolver: Arc<dyn WorkspaceResolver>) -> Self {
|
pub fn new(workspace: Arc<Workspace>) -> Self {
|
||||||
Self { resolver }
|
Self { workspace }
|
||||||
}
|
|
||||||
|
|
||||||
/// Create from a fixed workspace (backward compatibility).
|
|
||||||
pub fn from_workspace(workspace: Arc<Workspace>) -> Self {
|
|
||||||
Self {
|
|
||||||
resolver: Arc::new(FixedWorkspaceResolver::new(workspace)),
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -585,7 +456,7 @@ impl Tool for MemoryReadTool {
|
|||||||
async fn execute(
|
async fn execute(
|
||||||
&self,
|
&self,
|
||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
ctx: &JobContext,
|
_ctx: &JobContext,
|
||||||
) -> Result<ToolOutput, ToolError> {
|
) -> Result<ToolOutput, ToolError> {
|
||||||
let start = std::time::Instant::now();
|
let start = std::time::Instant::now();
|
||||||
|
|
||||||
@@ -599,8 +470,8 @@ impl Tool for MemoryReadTool {
|
|||||||
)));
|
)));
|
||||||
}
|
}
|
||||||
|
|
||||||
let workspace = self.resolver.resolve(&ctx.user_id).await;
|
let doc = self
|
||||||
let doc = workspace
|
.workspace
|
||||||
.read(path)
|
.read(path)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("Read failed: {}", e)))?;
|
.map_err(|e| ToolError::ExecutionFailed(format!("Read failed: {}", e)))?;
|
||||||
@@ -624,27 +495,20 @@ impl Tool for MemoryReadTool {
|
|||||||
///
|
///
|
||||||
/// Returns a hierarchical view of files and directories with configurable depth.
|
/// Returns a hierarchical view of files and directories with configurable depth.
|
||||||
pub struct MemoryTreeTool {
|
pub struct MemoryTreeTool {
|
||||||
resolver: Arc<dyn WorkspaceResolver>,
|
workspace: Arc<Workspace>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl MemoryTreeTool {
|
impl MemoryTreeTool {
|
||||||
/// Create a new memory tree tool with a workspace resolver.
|
/// Create a new memory tree tool.
|
||||||
pub fn new(resolver: Arc<dyn WorkspaceResolver>) -> Self {
|
pub fn new(workspace: Arc<Workspace>) -> Self {
|
||||||
Self { resolver }
|
Self { workspace }
|
||||||
}
|
|
||||||
|
|
||||||
/// Create from a fixed workspace (backward compatibility).
|
|
||||||
pub fn from_workspace(workspace: Arc<Workspace>) -> Self {
|
|
||||||
Self {
|
|
||||||
resolver: Arc::new(FixedWorkspaceResolver::new(workspace)),
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Recursively build tree structure.
|
/// Recursively build tree structure.
|
||||||
///
|
///
|
||||||
/// Returns a compact format where directories end with `/` and may have children.
|
/// Returns a compact format where directories end with `/` and may have children.
|
||||||
async fn build_tree(
|
async fn build_tree(
|
||||||
workspace: &Arc<Workspace>,
|
&self,
|
||||||
path: &str,
|
path: &str,
|
||||||
current_depth: usize,
|
current_depth: usize,
|
||||||
max_depth: usize,
|
max_depth: usize,
|
||||||
@@ -653,7 +517,8 @@ impl MemoryTreeTool {
|
|||||||
return Ok(Vec::new());
|
return Ok(Vec::new());
|
||||||
}
|
}
|
||||||
|
|
||||||
let entries = workspace
|
let entries = self
|
||||||
|
.workspace
|
||||||
.list(path)
|
.list(path)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("Tree failed: {}", e)))?;
|
.map_err(|e| ToolError::ExecutionFailed(format!("Tree failed: {}", e)))?;
|
||||||
@@ -668,13 +533,8 @@ impl MemoryTreeTool {
|
|||||||
};
|
};
|
||||||
|
|
||||||
if entry.is_directory && current_depth < max_depth {
|
if entry.is_directory && current_depth < max_depth {
|
||||||
let children = Box::pin(Self::build_tree(
|
let children =
|
||||||
workspace,
|
Box::pin(self.build_tree(&entry.path, current_depth + 1, max_depth)).await?;
|
||||||
&entry.path,
|
|
||||||
current_depth + 1,
|
|
||||||
max_depth,
|
|
||||||
))
|
|
||||||
.await?;
|
|
||||||
if children.is_empty() {
|
if children.is_empty() {
|
||||||
result.push(serde_json::Value::String(display_path));
|
result.push(serde_json::Value::String(display_path));
|
||||||
} else {
|
} else {
|
||||||
@@ -724,7 +584,7 @@ impl Tool for MemoryTreeTool {
|
|||||||
async fn execute(
|
async fn execute(
|
||||||
&self,
|
&self,
|
||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
ctx: &JobContext,
|
_ctx: &JobContext,
|
||||||
) -> Result<ToolOutput, ToolError> {
|
) -> Result<ToolOutput, ToolError> {
|
||||||
let start = std::time::Instant::now();
|
let start = std::time::Instant::now();
|
||||||
|
|
||||||
@@ -736,8 +596,7 @@ impl Tool for MemoryTreeTool {
|
|||||||
.unwrap_or(1)
|
.unwrap_or(1)
|
||||||
.clamp(1, 10) as usize;
|
.clamp(1, 10) as usize;
|
||||||
|
|
||||||
let workspace = self.resolver.resolve(&ctx.user_id).await;
|
let tree = self.build_tree(path, 1, depth).await?;
|
||||||
let tree = Self::build_tree(&workspace, path, 1, depth).await?;
|
|
||||||
|
|
||||||
// Compact output: just the tree array
|
// Compact output: just the tree array
|
||||||
Ok(ToolOutput::success(
|
Ok(ToolOutput::success(
|
||||||
@@ -791,7 +650,7 @@ mod tests {
|
|||||||
#[test]
|
#[test]
|
||||||
fn test_memory_search_schema() {
|
fn test_memory_search_schema() {
|
||||||
let workspace = make_test_workspace();
|
let workspace = make_test_workspace();
|
||||||
let tool = MemorySearchTool::from_workspace(workspace);
|
let tool = MemorySearchTool::new(workspace);
|
||||||
|
|
||||||
assert_eq!(tool.name(), "memory_search");
|
assert_eq!(tool.name(), "memory_search");
|
||||||
assert!(!tool.requires_sanitization());
|
assert!(!tool.requires_sanitization());
|
||||||
@@ -809,7 +668,7 @@ mod tests {
|
|||||||
#[test]
|
#[test]
|
||||||
fn test_memory_write_schema() {
|
fn test_memory_write_schema() {
|
||||||
let workspace = make_test_workspace();
|
let workspace = make_test_workspace();
|
||||||
let tool = MemoryWriteTool::from_workspace(workspace);
|
let tool = MemoryWriteTool::new(workspace);
|
||||||
|
|
||||||
assert_eq!(tool.name(), "memory_write");
|
assert_eq!(tool.name(), "memory_write");
|
||||||
|
|
||||||
@@ -822,7 +681,7 @@ mod tests {
|
|||||||
#[test]
|
#[test]
|
||||||
fn test_memory_read_schema() {
|
fn test_memory_read_schema() {
|
||||||
let workspace = make_test_workspace();
|
let workspace = make_test_workspace();
|
||||||
let tool = MemoryReadTool::from_workspace(workspace);
|
let tool = MemoryReadTool::new(workspace);
|
||||||
|
|
||||||
assert_eq!(tool.name(), "memory_read");
|
assert_eq!(tool.name(), "memory_read");
|
||||||
|
|
||||||
@@ -839,7 +698,7 @@ mod tests {
|
|||||||
#[test]
|
#[test]
|
||||||
fn test_memory_tree_schema() {
|
fn test_memory_tree_schema() {
|
||||||
let workspace = make_test_workspace();
|
let workspace = make_test_workspace();
|
||||||
let tool = MemoryTreeTool::from_workspace(workspace);
|
let tool = MemoryTreeTool::new(workspace);
|
||||||
|
|
||||||
assert_eq!(tool.name(), "memory_tree");
|
assert_eq!(tool.name(), "memory_tree");
|
||||||
|
|
||||||
@@ -852,7 +711,7 @@ mod tests {
|
|||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_memory_write_rejects_injection_to_identity_file() {
|
async fn test_memory_write_rejects_injection_to_identity_file() {
|
||||||
let workspace = make_test_workspace();
|
let workspace = make_test_workspace();
|
||||||
let tool = MemoryWriteTool::from_workspace(workspace);
|
let tool = MemoryWriteTool::new(workspace);
|
||||||
let ctx = JobContext::default();
|
let ctx = JobContext::default();
|
||||||
|
|
||||||
let params = serde_json::json!({
|
let params = serde_json::json!({
|
||||||
@@ -874,176 +733,4 @@ mod tests {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Regression tests for per-user workspace scoping (multi-tenant mode).
|
|
||||||
// See: https://github.com/nearai/ironclaw/pull/1118
|
|
||||||
// Bug: memory tools used a single startup workspace regardless of which
|
|
||||||
// user was chatting. Fix: resolve workspace per-request via JobContext.user_id.
|
|
||||||
|
|
||||||
#[cfg(feature = "postgres")]
|
|
||||||
mod resolver_tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
fn make_test_workspace_for_user(user_id: &str) -> Arc<Workspace> {
|
|
||||||
Arc::new(Workspace::new(
|
|
||||||
user_id,
|
|
||||||
deadpool_postgres::Pool::builder(deadpool_postgres::Manager::new(
|
|
||||||
tokio_postgres::Config::new(),
|
|
||||||
tokio_postgres::NoTls,
|
|
||||||
))
|
|
||||||
.build()
|
|
||||||
.unwrap(),
|
|
||||||
))
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_fixed_workspace_resolver_ignores_user_id() {
|
|
||||||
let ws = make_test_workspace_for_user("alice");
|
|
||||||
let resolver = FixedWorkspaceResolver::new(Arc::clone(&ws));
|
|
||||||
|
|
||||||
let ws_alice = resolver.resolve("alice").await;
|
|
||||||
let ws_bob = resolver.resolve("bob").await;
|
|
||||||
|
|
||||||
// Both should return the exact same Arc (pointer equality)
|
|
||||||
assert!(Arc::ptr_eq(&ws_alice, &ws_bob));
|
|
||||||
assert_eq!(ws_alice.user_id(), "alice");
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Tracking resolver that records which user_ids were requested.
|
|
||||||
struct TrackingWorkspaceResolver {
|
|
||||||
inner: FixedWorkspaceResolver,
|
|
||||||
resolved_users: std::sync::Mutex<Vec<String>>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl TrackingWorkspaceResolver {
|
|
||||||
fn new(workspace: Arc<Workspace>) -> Self {
|
|
||||||
Self {
|
|
||||||
inner: FixedWorkspaceResolver::new(workspace),
|
|
||||||
resolved_users: std::sync::Mutex::new(Vec::new()),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn resolved_users(&self) -> Vec<String> {
|
|
||||||
self.resolved_users.lock().unwrap().clone()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[async_trait]
|
|
||||||
impl WorkspaceResolver for TrackingWorkspaceResolver {
|
|
||||||
async fn resolve(&self, user_id: &str) -> Arc<Workspace> {
|
|
||||||
self.resolved_users
|
|
||||||
.lock()
|
|
||||||
.unwrap()
|
|
||||||
.push(user_id.to_string());
|
|
||||||
self.inner.resolve(user_id).await
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_memory_search_uses_job_context_user_id() {
|
|
||||||
let ws = make_test_workspace_for_user("default");
|
|
||||||
let tracker = Arc::new(TrackingWorkspaceResolver::new(ws));
|
|
||||||
let tool = MemorySearchTool::new(tracker.clone() as Arc<dyn WorkspaceResolver>);
|
|
||||||
|
|
||||||
// Execute with user_id "alice"
|
|
||||||
let ctx_alice = JobContext::with_user("alice", "test", "test");
|
|
||||||
let params = serde_json::json!({"query": "test"});
|
|
||||||
// The search will fail (no real DB) but we only care about resolver call
|
|
||||||
let _ = tool.execute(params, &ctx_alice).await;
|
|
||||||
|
|
||||||
// Execute with user_id "bob"
|
|
||||||
let ctx_bob = JobContext::with_user("bob", "test", "test");
|
|
||||||
let params = serde_json::json!({"query": "test"});
|
|
||||||
let _ = tool.execute(params, &ctx_bob).await;
|
|
||||||
|
|
||||||
let resolved = tracker.resolved_users();
|
|
||||||
assert_eq!(resolved, vec!["alice", "bob"]);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_memory_write_uses_job_context_user_id() {
|
|
||||||
let ws = make_test_workspace_for_user("default");
|
|
||||||
let tracker = Arc::new(TrackingWorkspaceResolver::new(ws));
|
|
||||||
let tool = MemoryWriteTool::new(tracker.clone() as Arc<dyn WorkspaceResolver>);
|
|
||||||
|
|
||||||
// Execute with user_id "alice"
|
|
||||||
let ctx_alice = JobContext::with_user("alice", "test", "test");
|
|
||||||
let params = serde_json::json!({
|
|
||||||
"content": "remember this",
|
|
||||||
"target": "daily_log",
|
|
||||||
});
|
|
||||||
let _ = tool.execute(params, &ctx_alice).await;
|
|
||||||
|
|
||||||
// Execute with user_id "bob"
|
|
||||||
let ctx_bob = JobContext::with_user("bob", "test", "test");
|
|
||||||
let params = serde_json::json!({
|
|
||||||
"content": "remember that",
|
|
||||||
"target": "daily_log",
|
|
||||||
});
|
|
||||||
let _ = tool.execute(params, &ctx_bob).await;
|
|
||||||
|
|
||||||
let resolved = tracker.resolved_users();
|
|
||||||
assert_eq!(resolved, vec!["alice", "bob"]);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(feature = "libsql")]
|
|
||||||
mod per_user_resolver_tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
async fn make_test_db() -> Arc<dyn crate::db::Database> {
|
|
||||||
use crate::db::libsql::LibSqlBackend;
|
|
||||||
let temp_dir = tempfile::tempdir().expect("tempdir");
|
|
||||||
let db_path = temp_dir.path().join("resolver_test.db");
|
|
||||||
let backend = LibSqlBackend::new_local(&db_path)
|
|
||||||
.await
|
|
||||||
.expect("LibSqlBackend");
|
|
||||||
<LibSqlBackend as crate::db::Database>::run_migrations(&backend)
|
|
||||||
.await
|
|
||||||
.expect("migrations");
|
|
||||||
// Leak the tempdir so it outlives the test (cleaned up on process exit).
|
|
||||||
std::mem::forget(temp_dir);
|
|
||||||
Arc::new(backend)
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_per_user_workspace_resolver_returns_different_workspaces() {
|
|
||||||
let db = make_test_db().await;
|
|
||||||
|
|
||||||
let resolver = PerUserWorkspaceResolver::new(
|
|
||||||
db,
|
|
||||||
None,
|
|
||||||
crate::workspace::EmbeddingCacheConfig::default(),
|
|
||||||
crate::config::WorkspaceSearchConfig::default(),
|
|
||||||
crate::config::WorkspaceConfig::default(),
|
|
||||||
);
|
|
||||||
|
|
||||||
let ws_alice = resolver.resolve("alice").await;
|
|
||||||
let ws_bob = resolver.resolve("bob").await;
|
|
||||||
|
|
||||||
// Different user IDs should get different workspaces
|
|
||||||
assert_eq!(ws_alice.user_id(), "alice");
|
|
||||||
assert_eq!(ws_bob.user_id(), "bob");
|
|
||||||
assert!(!Arc::ptr_eq(&ws_alice, &ws_bob));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn test_per_user_workspace_resolver_caches_workspace() {
|
|
||||||
let db = make_test_db().await;
|
|
||||||
|
|
||||||
let resolver = PerUserWorkspaceResolver::new(
|
|
||||||
db,
|
|
||||||
None,
|
|
||||||
crate::workspace::EmbeddingCacheConfig::default(),
|
|
||||||
crate::config::WorkspaceSearchConfig::default(),
|
|
||||||
crate::config::WorkspaceConfig::default(),
|
|
||||||
);
|
|
||||||
|
|
||||||
let ws1 = resolver.resolve("alice").await;
|
|
||||||
let ws2 = resolver.resolve("alice").await;
|
|
||||||
|
|
||||||
// Same user_id should return the same cached Arc (pointer equality)
|
|
||||||
assert!(Arc::ptr_eq(&ws1, &ws2));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ mod file;
|
|||||||
mod http;
|
mod http;
|
||||||
mod job;
|
mod job;
|
||||||
mod json;
|
mod json;
|
||||||
pub mod memory;
|
mod memory;
|
||||||
mod message;
|
mod message;
|
||||||
pub mod path_utils;
|
pub mod path_utils;
|
||||||
mod restart;
|
mod restart;
|
||||||
|
|||||||
+6
-32
@@ -334,37 +334,15 @@ impl ToolRegistry {
|
|||||||
tracing::debug!("Registered 5 development tools");
|
tracing::debug!("Registered 5 development tools");
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Register memory tools with a workspace resolver.
|
/// Register memory tools with a workspace.
|
||||||
///
|
|
||||||
/// Memory tools require a workspace resolver for persistence. Call this after
|
|
||||||
/// `register_builtin_tools()` if you have a workspace available.
|
|
||||||
pub fn register_memory_tools_with_resolver(
|
|
||||||
&self,
|
|
||||||
resolver: Arc<dyn crate::tools::builtin::memory::WorkspaceResolver>,
|
|
||||||
) {
|
|
||||||
self.register_sync(Arc::new(MemorySearchTool::new(Arc::clone(&resolver))));
|
|
||||||
self.register_sync(Arc::new(MemoryWriteTool::new(Arc::clone(&resolver))));
|
|
||||||
self.register_sync(Arc::new(MemoryReadTool::new(Arc::clone(&resolver))));
|
|
||||||
self.register_sync(Arc::new(MemoryTreeTool::new(resolver)));
|
|
||||||
|
|
||||||
tracing::debug!("Registered 4 memory tools");
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Register memory tools with a fixed workspace (backward compatibility).
|
|
||||||
///
|
///
|
||||||
/// Memory tools require a workspace for persistence. Call this after
|
/// Memory tools require a workspace for persistence. Call this after
|
||||||
/// `register_builtin_tools()` if you have a workspace available.
|
/// `register_builtin_tools()` if you have a workspace available.
|
||||||
pub fn register_memory_tools(&self, workspace: Arc<Workspace>) {
|
pub fn register_memory_tools(&self, workspace: Arc<Workspace>) {
|
||||||
self.register_sync(Arc::new(MemorySearchTool::from_workspace(Arc::clone(
|
self.register_sync(Arc::new(MemorySearchTool::new(Arc::clone(&workspace))));
|
||||||
&workspace,
|
self.register_sync(Arc::new(MemoryWriteTool::new(Arc::clone(&workspace))));
|
||||||
))));
|
self.register_sync(Arc::new(MemoryReadTool::new(Arc::clone(&workspace))));
|
||||||
self.register_sync(Arc::new(MemoryWriteTool::from_workspace(Arc::clone(
|
self.register_sync(Arc::new(MemoryTreeTool::new(workspace)));
|
||||||
&workspace,
|
|
||||||
))));
|
|
||||||
self.register_sync(Arc::new(MemoryReadTool::from_workspace(Arc::clone(
|
|
||||||
&workspace,
|
|
||||||
))));
|
|
||||||
self.register_sync(Arc::new(MemoryTreeTool::from_workspace(workspace)));
|
|
||||||
|
|
||||||
tracing::debug!("Registered 4 memory tools");
|
tracing::debug!("Registered 4 memory tools");
|
||||||
}
|
}
|
||||||
@@ -383,11 +361,7 @@ impl ToolRegistry {
|
|||||||
job_manager: Option<Arc<ContainerJobManager>>,
|
job_manager: Option<Arc<ContainerJobManager>>,
|
||||||
store: Option<Arc<dyn Database>>,
|
store: Option<Arc<dyn Database>>,
|
||||||
job_event_tx: Option<
|
job_event_tx: Option<
|
||||||
tokio::sync::broadcast::Sender<(
|
tokio::sync::broadcast::Sender<(uuid::Uuid, crate::channels::web::types::SseEvent)>,
|
||||||
uuid::Uuid,
|
|
||||||
String,
|
|
||||||
crate::channels::web::types::SseEvent,
|
|
||||||
)>,
|
|
||||||
>,
|
>,
|
||||||
inject_tx: Option<tokio::sync::mpsc::Sender<crate::channels::IncomingMessage>>,
|
inject_tx: Option<tokio::sync::mpsc::Sender<crate::channels::IncomingMessage>>,
|
||||||
prompt_queue: Option<PromptQueue>,
|
prompt_queue: Option<PromptQueue>,
|
||||||
|
|||||||
@@ -1,249 +0,0 @@
|
|||||||
"""OAuth URL parameter validation e2e tests.
|
|
||||||
|
|
||||||
Tests for bug #992: Google OAuth URL broken when initiated from Telegram.
|
|
||||||
Specifically verifies that OAuth query parameters are correctly formatted:
|
|
||||||
- "client_id" (with underscore) NOT "clientid" (without underscore)
|
|
||||||
- All standard OAuth parameters are present and correctly encoded
|
|
||||||
- URLs are consistent across channels (web, Telegram, etc.)
|
|
||||||
|
|
||||||
The test verifies:
|
|
||||||
1. OAuth URL is generated with correct parameters
|
|
||||||
2. URL works with the OAuth provider (Google)
|
|
||||||
3. Extra parameters (access_type, prompt) are preserved
|
|
||||||
"""
|
|
||||||
|
|
||||||
from urllib.parse import parse_qs, urlparse
|
|
||||||
import pytest
|
|
||||||
|
|
||||||
from helpers import api_post, api_get
|
|
||||||
|
|
||||||
|
|
||||||
async def _extract_oauth_params(auth_url: str) -> dict:
|
|
||||||
"""Extract and validate OAuth query parameters from auth_url.
|
|
||||||
|
|
||||||
Returns dict with parsed parameters:
|
|
||||||
{
|
|
||||||
'client_id': '...',
|
|
||||||
'redirect_uri': '...',
|
|
||||||
'response_type': 'code',
|
|
||||||
'scope': '...',
|
|
||||||
'state': '...',
|
|
||||||
'access_type': '...',
|
|
||||||
'prompt': '...',
|
|
||||||
...
|
|
||||||
}
|
|
||||||
"""
|
|
||||||
parsed = urlparse(auth_url)
|
|
||||||
qs = parse_qs(parsed.query)
|
|
||||||
|
|
||||||
# Convert lists to single values for easier testing
|
|
||||||
params = {k: v[0] if len(v) > 0 else v for k, v in qs.items()}
|
|
||||||
return params
|
|
||||||
|
|
||||||
|
|
||||||
async def _get_extension(ironclaw_server, name):
|
|
||||||
"""Get a specific extension from the extensions list, or None."""
|
|
||||||
r = await api_get(ironclaw_server, "/api/extensions")
|
|
||||||
for ext in r.json().get("extensions", []):
|
|
||||||
if ext["name"] == name:
|
|
||||||
return ext
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
|
||||||
async def installed_gmail(ironclaw_server):
|
|
||||||
"""Installs the 'gmail' extension before a test and removes it after.
|
|
||||||
|
|
||||||
This fixture handles the setup and teardown of the Gmail extension,
|
|
||||||
ensuring a clean state for each test.
|
|
||||||
"""
|
|
||||||
# Ensure Gmail is not installed before test
|
|
||||||
ext = await _get_extension(ironclaw_server, "gmail")
|
|
||||||
if ext:
|
|
||||||
r = await api_post(ironclaw_server, "/api/extensions/gmail/remove", timeout=30)
|
|
||||||
assert r.status_code == 200
|
|
||||||
|
|
||||||
# Install Gmail
|
|
||||||
r = await api_post(
|
|
||||||
ironclaw_server,
|
|
||||||
"/api/extensions/install",
|
|
||||||
json={"name": "gmail"},
|
|
||||||
timeout=180,
|
|
||||||
)
|
|
||||||
assert r.status_code == 200, f"Gmail install failed: {r.text}"
|
|
||||||
assert r.json().get("success") is True, f"Install failed: {r.json().get('message', '')}"
|
|
||||||
|
|
||||||
yield
|
|
||||||
|
|
||||||
# Teardown: remove gmail
|
|
||||||
r = await api_post(ironclaw_server, "/api/extensions/gmail/remove", timeout=30)
|
|
||||||
assert r.status_code == 200, f"Gmail removal failed: {r.text}"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
|
||||||
async def auth_url(ironclaw_server, installed_gmail):
|
|
||||||
"""Generate and return an OAuth auth URL.
|
|
||||||
|
|
||||||
Requires Gmail to be installed (depends on installed_gmail fixture).
|
|
||||||
"""
|
|
||||||
r = await api_post(
|
|
||||||
ironclaw_server,
|
|
||||||
"/api/extensions/gmail/setup",
|
|
||||||
json={"secrets": {}},
|
|
||||||
timeout=30,
|
|
||||||
)
|
|
||||||
assert r.status_code == 200
|
|
||||||
data = r.json()
|
|
||||||
assert data.get("success") is True, f"Setup failed: {data.get('message', '')}"
|
|
||||||
|
|
||||||
url = data.get("auth_url")
|
|
||||||
assert url is not None, f"Expected auth_url in response: {data}"
|
|
||||||
assert "accounts.google.com" in url, f"auth_url should point to Google: {url}"
|
|
||||||
|
|
||||||
return url
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
|
||||||
async def oauth_params(auth_url):
|
|
||||||
"""Extract and return OAuth parameters from auth_url.
|
|
||||||
|
|
||||||
Depends on auth_url fixture.
|
|
||||||
"""
|
|
||||||
return await _extract_oauth_params(auth_url)
|
|
||||||
|
|
||||||
|
|
||||||
# ─ OAuth URL parameter validation tests ────────────────────────────────
|
|
||||||
|
|
||||||
async def test_oauth_url_has_client_id_not_clientid(oauth_params, auth_url):
|
|
||||||
"""Verify OAuth URL has 'client_id' (with underscore), NOT 'clientid'.
|
|
||||||
|
|
||||||
Bug #992: Ensure the parameter name is correct across all channels.
|
|
||||||
"""
|
|
||||||
params = oauth_params
|
|
||||||
|
|
||||||
# The bug: "clientid" appears instead of "client_id"
|
|
||||||
# Verify the CORRECT parameter name exists
|
|
||||||
assert "client_id" in params, (
|
|
||||||
f"OAuth URL missing 'client_id' parameter. "
|
|
||||||
f"URL: {auth_url}\nParams: {params}"
|
|
||||||
)
|
|
||||||
assert params["client_id"], "client_id should have a value"
|
|
||||||
|
|
||||||
# Verify the INCORRECT parameter name does NOT exist
|
|
||||||
assert "clientid" not in params, (
|
|
||||||
f"OAuth URL should NOT have 'clientid' (without underscore). "
|
|
||||||
f"Bug #992: URL: {auth_url}\nParams: {params}"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
async def test_oauth_url_has_required_parameters(oauth_params):
|
|
||||||
"""Verify all required OAuth 2.0 parameters are present."""
|
|
||||||
params = oauth_params
|
|
||||||
|
|
||||||
# Required OAuth 2.0 parameters
|
|
||||||
required = ["client_id", "response_type", "redirect_uri", "scope", "state"]
|
|
||||||
for param in required:
|
|
||||||
assert param in params, (
|
|
||||||
f"Missing required OAuth parameter: {param}. "
|
|
||||||
f"Params: {params}"
|
|
||||||
)
|
|
||||||
assert params[param], f"Parameter '{param}' should have a non-empty value"
|
|
||||||
|
|
||||||
# Validate specific values
|
|
||||||
assert params["response_type"] == "code", "Should use authorization_code flow"
|
|
||||||
assert "oauth" in params["redirect_uri"], "Redirect URI should be an OAuth callback"
|
|
||||||
|
|
||||||
|
|
||||||
async def test_oauth_url_has_extra_params(oauth_params):
|
|
||||||
"""Verify extra_params from capabilities.json are included."""
|
|
||||||
params = oauth_params
|
|
||||||
|
|
||||||
# Google-specific extra_params from gmail-tool.capabilities.json
|
|
||||||
assert "access_type" in params, (
|
|
||||||
"Should include 'access_type' from extra_params"
|
|
||||||
)
|
|
||||||
assert params["access_type"] == "offline", (
|
|
||||||
"access_type should be 'offline' for Gmail"
|
|
||||||
)
|
|
||||||
|
|
||||||
assert "prompt" in params, (
|
|
||||||
"Should include 'prompt' from extra_params"
|
|
||||||
)
|
|
||||||
assert params["prompt"] == "consent", (
|
|
||||||
"prompt should be 'consent' for Gmail"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
async def test_oauth_url_is_valid_google_oauth(auth_url):
|
|
||||||
"""Verify the URL is a valid Google OAuth 2.0 authorization URL."""
|
|
||||||
# Verify scheme and host
|
|
||||||
parsed = urlparse(auth_url)
|
|
||||||
assert parsed.scheme == "https", "OAuth URL must use HTTPS"
|
|
||||||
assert "accounts.google.com" in parsed.netloc, "Must be Google's OAuth endpoint"
|
|
||||||
assert parsed.path == "/o/oauth2/v2/auth", "Must use Google OAuth 2.0 endpoint"
|
|
||||||
|
|
||||||
|
|
||||||
async def test_oauth_url_state_is_unique(ironclaw_server, installed_gmail, oauth_params, auth_url):
|
|
||||||
"""Verify CSRF state is present and unique per request."""
|
|
||||||
# Get a new OAuth URL
|
|
||||||
r = await api_post(
|
|
||||||
ironclaw_server,
|
|
||||||
"/api/extensions/gmail/setup",
|
|
||||||
json={"secrets": {}},
|
|
||||||
timeout=30,
|
|
||||||
)
|
|
||||||
assert r.status_code == 200
|
|
||||||
new_auth_url = r.json().get("auth_url")
|
|
||||||
assert new_auth_url is not None
|
|
||||||
|
|
||||||
# Extract state from both URLs
|
|
||||||
original_params = oauth_params
|
|
||||||
new_params = await _extract_oauth_params(new_auth_url)
|
|
||||||
|
|
||||||
original_state = original_params.get("state")
|
|
||||||
new_state = new_params.get("state")
|
|
||||||
|
|
||||||
assert original_state is not None, "Should have state parameter"
|
|
||||||
assert new_state is not None, "New request should have state parameter"
|
|
||||||
assert original_state != new_state, (
|
|
||||||
"CSRF state should be unique per request (for security)"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
async def test_oauth_url_escaping(auth_url):
|
|
||||||
"""Verify URL query parameters are properly escaped."""
|
|
||||||
# Verify special characters in values are URL-encoded
|
|
||||||
# For example, scopes contain spaces which should be %20
|
|
||||||
assert "%20" in auth_url or "+" in auth_url or "%2B" in auth_url or " " not in auth_url, (
|
|
||||||
"OAuth URL should properly encode special characters in parameters"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
# ─ Telegram-specific tests (when Telegram channel is available) ──────────
|
|
||||||
|
|
||||||
class TestOAuthURLViaTelegram:
|
|
||||||
"""Test OAuth URL generation specifically via Telegram channel.
|
|
||||||
|
|
||||||
These tests would verify that the same OAuth URL works correctly when
|
|
||||||
transmitted through the Telegram WASM channel (as opposed to web gateway).
|
|
||||||
|
|
||||||
Currently marked as xfail pending Telegram channel setup in E2E tests.
|
|
||||||
"""
|
|
||||||
|
|
||||||
@pytest.mark.skip(reason="Telegram channel E2E setup not yet implemented")
|
|
||||||
async def test_telegram_oauth_url_has_correct_parameters(self):
|
|
||||||
"""Verify OAuth URL sent via Telegram has correct parameter names."""
|
|
||||||
# This test would:
|
|
||||||
# 1. Send a message via Telegram that triggers OAuth
|
|
||||||
# 2. Capture the status update sent to Telegram
|
|
||||||
# 3. Extract the auth_url from the message
|
|
||||||
# 4. Verify it has "client_id" not "clientid"
|
|
||||||
pass
|
|
||||||
|
|
||||||
@pytest.mark.skip(reason="Telegram channel E2E setup not yet implemented")
|
|
||||||
async def test_telegram_oauth_url_can_be_regenerated(self):
|
|
||||||
"""Verify OAuth URL can be regenerated when requested via Telegram."""
|
|
||||||
# This test would verify that the bug #992 symptom
|
|
||||||
# "URL cannot be regenerated when asked" is fixed.
|
|
||||||
# If the URL is cached incorrectly, regeneration would fail.
|
|
||||||
pass
|
|
||||||
@@ -661,7 +661,7 @@ mod advanced {
|
|||||||
.await
|
.await
|
||||||
.expect("failed to inject test token");
|
.expect("failed to inject test token");
|
||||||
|
|
||||||
let activate_result = ext_mgr.activate("mock-notion", "default").await;
|
let activate_result = ext_mgr.activate("mock-notion").await;
|
||||||
assert!(
|
assert!(
|
||||||
activate_result.is_ok(),
|
activate_result.is_ok(),
|
||||||
"activation failed: {:?}",
|
"activation failed: {:?}",
|
||||||
|
|||||||
@@ -216,7 +216,7 @@ async fn extension_manager_with_process_manager_constructs() {
|
|||||||
);
|
);
|
||||||
|
|
||||||
// Verify the manager is functional — list returns Ok.
|
// Verify the manager is functional — list returns Ok.
|
||||||
let result = manager.list(None, false, "test").await;
|
let result = manager.list(None, false).await;
|
||||||
assert!(result.is_ok(), "list should succeed on empty manager");
|
assert!(result.is_ok(), "list should succeed on empty manager");
|
||||||
assert!(result.unwrap().is_empty());
|
assert!(result.unwrap().is_empty());
|
||||||
}
|
}
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -1,240 +0,0 @@
|
|||||||
//! Tests proving that multi-tenant system prompts are broken.
|
|
||||||
//!
|
|
||||||
//! Bug: In multi-tenant mode, the agent loop uses `self.workspace()` which
|
|
||||||
//! returns a single shared workspace (user_id="default"). Identity files
|
|
||||||
//! (IDENTITY.md, SOUL.md, USER.md) seeded under per-user IDs ("alice",
|
|
||||||
//! "bob") are invisible to this workspace, so the system prompt is
|
|
||||||
//! empty/wrong.
|
|
||||||
//!
|
|
||||||
//! These tests:
|
|
||||||
//! 1. Seed identity files for two users (alice, bob) in the database
|
|
||||||
//! 2. Send messages as each user
|
|
||||||
//! 3. Verify the system prompt in captured LLM requests contains the
|
|
||||||
//! correct user's identity
|
|
||||||
//! 4. Verify user A's identity doesn't leak into user B's prompt
|
|
||||||
//!
|
|
||||||
//! All tests are expected to FAIL until the bug is fixed.
|
|
||||||
|
|
||||||
#[cfg(feature = "libsql")]
|
|
||||||
mod support;
|
|
||||||
|
|
||||||
#[cfg(feature = "libsql")]
|
|
||||||
mod tests {
|
|
||||||
use std::sync::Arc;
|
|
||||||
use std::time::Duration;
|
|
||||||
|
|
||||||
use ironclaw::channels::IncomingMessage;
|
|
||||||
use ironclaw::llm::Role;
|
|
||||||
use ironclaw::workspace::Workspace;
|
|
||||||
|
|
||||||
use crate::support::test_rig::TestRigBuilder;
|
|
||||||
use crate::support::trace_llm::{LlmTrace, TraceResponse, TraceStep};
|
|
||||||
|
|
||||||
const TIMEOUT: Duration = Duration::from_secs(15);
|
|
||||||
|
|
||||||
const ALICE_USER_ID: &str = "alice";
|
|
||||||
const BOB_USER_ID: &str = "bob";
|
|
||||||
|
|
||||||
const ALICE_IDENTITY: &str = "You are Alice's personal assistant. \
|
|
||||||
Alice is a software engineer who lives in Seattle.";
|
|
||||||
const BOB_IDENTITY: &str = "You are Bob's personal assistant. \
|
|
||||||
Bob is a marine biologist who lives in Miami.";
|
|
||||||
|
|
||||||
/// Create a simple trace that returns a canned text response.
|
|
||||||
/// We need one step per message we plan to send.
|
|
||||||
fn simple_trace(num_steps: usize) -> LlmTrace {
|
|
||||||
let steps: Vec<TraceStep> = (0..num_steps)
|
|
||||||
.map(|i| TraceStep {
|
|
||||||
request_hint: None,
|
|
||||||
response: TraceResponse::Text {
|
|
||||||
content: format!("Response {}", i),
|
|
||||||
input_tokens: 100,
|
|
||||||
output_tokens: 10,
|
|
||||||
},
|
|
||||||
expected_tool_results: Vec::new(),
|
|
||||||
})
|
|
||||||
.collect();
|
|
||||||
|
|
||||||
// Create separate turns for each step so the trace replays correctly.
|
|
||||||
let turns: Vec<crate::support::trace_llm::TraceTurn> = steps
|
|
||||||
.into_iter()
|
|
||||||
.enumerate()
|
|
||||||
.map(|(i, step)| crate::support::trace_llm::TraceTurn {
|
|
||||||
user_input: format!("message {}", i),
|
|
||||||
steps: vec![step],
|
|
||||||
expects: Default::default(),
|
|
||||||
})
|
|
||||||
.collect();
|
|
||||||
|
|
||||||
LlmTrace::new("test-model", turns)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Seed identity files for a user by creating a workspace scoped to that
|
|
||||||
/// user and writing IDENTITY.md.
|
|
||||||
async fn seed_identity(db: &Arc<dyn ironclaw::db::Database>, user_id: &str, content: &str) {
|
|
||||||
let ws = Workspace::new_with_db(user_id, db.clone());
|
|
||||||
ws.write("IDENTITY.md", content)
|
|
||||||
.await
|
|
||||||
.unwrap_or_else(|e| panic!("Failed to seed IDENTITY.md for {user_id}: {e}"));
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Extract the system prompt from captured LLM requests.
|
|
||||||
///
|
|
||||||
/// The system prompt is the first message with role=System in the first
|
|
||||||
/// LLM request for a given turn.
|
|
||||||
fn extract_system_prompt(requests: &[Vec<ironclaw::llm::ChatMessage>]) -> Option<String> {
|
|
||||||
requests.last().and_then(|msgs| {
|
|
||||||
msgs.iter()
|
|
||||||
.find(|m| matches!(m.role, Role::System))
|
|
||||||
.map(|m| m.content.clone())
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// -----------------------------------------------------------------------
|
|
||||||
// Test 1: Alice's identity should appear in system prompt when messaging
|
|
||||||
// as Alice.
|
|
||||||
// -----------------------------------------------------------------------
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn alice_system_prompt_contains_alice_identity() {
|
|
||||||
let trace = simple_trace(1);
|
|
||||||
let rig = TestRigBuilder::new().with_trace(trace).build().await;
|
|
||||||
|
|
||||||
// Seed alice's identity into the database
|
|
||||||
let db = rig.database();
|
|
||||||
seed_identity(db, ALICE_USER_ID, ALICE_IDENTITY).await;
|
|
||||||
|
|
||||||
// Send a message AS alice (using her user_id)
|
|
||||||
let msg = IncomingMessage::new("test", ALICE_USER_ID, "Hello, who am I?");
|
|
||||||
rig.send_incoming(msg).await;
|
|
||||||
let _responses = rig.wait_for_responses(1, TIMEOUT).await;
|
|
||||||
|
|
||||||
// The system prompt sent to the LLM should contain Alice's identity
|
|
||||||
let requests = rig.captured_llm_requests();
|
|
||||||
let system_prompt =
|
|
||||||
extract_system_prompt(&requests).expect("Expected a system prompt in the LLM request");
|
|
||||||
|
|
||||||
assert!(
|
|
||||||
system_prompt.contains("Alice is a software engineer"),
|
|
||||||
"System prompt should contain Alice's identity when messaging as Alice.\n\
|
|
||||||
Actual system prompt:\n{system_prompt}"
|
|
||||||
);
|
|
||||||
|
|
||||||
rig.shutdown();
|
|
||||||
}
|
|
||||||
|
|
||||||
// -----------------------------------------------------------------------
|
|
||||||
// Test 2: Bob's identity should appear in system prompt when messaging
|
|
||||||
// as Bob.
|
|
||||||
// -----------------------------------------------------------------------
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn bob_system_prompt_contains_bob_identity() {
|
|
||||||
let trace = simple_trace(1);
|
|
||||||
let rig = TestRigBuilder::new().with_trace(trace).build().await;
|
|
||||||
|
|
||||||
// Seed bob's identity into the database
|
|
||||||
let db = rig.database();
|
|
||||||
seed_identity(db, BOB_USER_ID, BOB_IDENTITY).await;
|
|
||||||
|
|
||||||
// Send a message AS bob
|
|
||||||
let msg = IncomingMessage::new("test", BOB_USER_ID, "Hello, who am I?");
|
|
||||||
rig.send_incoming(msg).await;
|
|
||||||
let _responses = rig.wait_for_responses(1, TIMEOUT).await;
|
|
||||||
|
|
||||||
// The system prompt should contain Bob's identity
|
|
||||||
let requests = rig.captured_llm_requests();
|
|
||||||
let system_prompt =
|
|
||||||
extract_system_prompt(&requests).expect("Expected a system prompt in the LLM request");
|
|
||||||
|
|
||||||
assert!(
|
|
||||||
system_prompt.contains("Bob is a marine biologist"),
|
|
||||||
"System prompt should contain Bob's identity when messaging as Bob.\n\
|
|
||||||
Actual system prompt:\n{system_prompt}"
|
|
||||||
);
|
|
||||||
|
|
||||||
rig.shutdown();
|
|
||||||
}
|
|
||||||
|
|
||||||
// -----------------------------------------------------------------------
|
|
||||||
// Test 3: Alice's identity must NOT appear in Bob's system prompt.
|
|
||||||
// -----------------------------------------------------------------------
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn alice_identity_does_not_leak_into_bob_prompt() {
|
|
||||||
let trace = simple_trace(1);
|
|
||||||
let rig = TestRigBuilder::new().with_trace(trace).build().await;
|
|
||||||
|
|
||||||
// Seed BOTH users' identities
|
|
||||||
let db = rig.database();
|
|
||||||
seed_identity(db, ALICE_USER_ID, ALICE_IDENTITY).await;
|
|
||||||
seed_identity(db, BOB_USER_ID, BOB_IDENTITY).await;
|
|
||||||
|
|
||||||
// Send a message AS bob
|
|
||||||
let msg = IncomingMessage::new("test", BOB_USER_ID, "Tell me about myself");
|
|
||||||
rig.send_incoming(msg).await;
|
|
||||||
let _responses = rig.wait_for_responses(1, TIMEOUT).await;
|
|
||||||
|
|
||||||
// Bob's prompt must NOT contain Alice's identity
|
|
||||||
let requests = rig.captured_llm_requests();
|
|
||||||
let system_prompt = extract_system_prompt(&requests);
|
|
||||||
|
|
||||||
if let Some(ref prompt) = system_prompt {
|
|
||||||
assert!(
|
|
||||||
!prompt.contains("Alice is a software engineer"),
|
|
||||||
"Alice's identity LEAKED into Bob's system prompt!\n\
|
|
||||||
System prompt:\n{prompt}"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
// Also verify Bob's identity IS present (compound check)
|
|
||||||
let prompt = system_prompt.expect("Expected a system prompt in the LLM request");
|
|
||||||
assert!(
|
|
||||||
prompt.contains("Bob is a marine biologist"),
|
|
||||||
"Bob's own identity should be in his system prompt.\n\
|
|
||||||
Actual system prompt:\n{prompt}"
|
|
||||||
);
|
|
||||||
|
|
||||||
rig.shutdown();
|
|
||||||
}
|
|
||||||
|
|
||||||
// -----------------------------------------------------------------------
|
|
||||||
// Test 4: Bob's identity must NOT appear in Alice's system prompt.
|
|
||||||
// -----------------------------------------------------------------------
|
|
||||||
|
|
||||||
#[tokio::test]
|
|
||||||
async fn bob_identity_does_not_leak_into_alice_prompt() {
|
|
||||||
let trace = simple_trace(1);
|
|
||||||
let rig = TestRigBuilder::new().with_trace(trace).build().await;
|
|
||||||
|
|
||||||
// Seed BOTH users' identities
|
|
||||||
let db = rig.database();
|
|
||||||
seed_identity(db, ALICE_USER_ID, ALICE_IDENTITY).await;
|
|
||||||
seed_identity(db, BOB_USER_ID, BOB_IDENTITY).await;
|
|
||||||
|
|
||||||
// Send a message AS alice
|
|
||||||
let msg = IncomingMessage::new("test", ALICE_USER_ID, "Tell me about myself");
|
|
||||||
rig.send_incoming(msg).await;
|
|
||||||
let _responses = rig.wait_for_responses(1, TIMEOUT).await;
|
|
||||||
|
|
||||||
// Alice's prompt must NOT contain Bob's identity
|
|
||||||
let requests = rig.captured_llm_requests();
|
|
||||||
let system_prompt = extract_system_prompt(&requests);
|
|
||||||
|
|
||||||
if let Some(ref prompt) = system_prompt {
|
|
||||||
assert!(
|
|
||||||
!prompt.contains("Bob is a marine biologist"),
|
|
||||||
"Bob's identity LEAKED into Alice's system prompt!\n\
|
|
||||||
System prompt:\n{prompt}"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
// Also verify Alice's identity IS present
|
|
||||||
let prompt = system_prompt.expect("Expected a system prompt in the LLM request");
|
|
||||||
assert!(
|
|
||||||
prompt.contains("Alice is a software engineer"),
|
|
||||||
"Alice's own identity should be in her system prompt.\n\
|
|
||||||
Actual system prompt:\n{prompt}"
|
|
||||||
);
|
|
||||||
|
|
||||||
rig.shutdown();
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -191,9 +191,8 @@ async fn start_test_server_with_provider(
|
|||||||
) -> (SocketAddr, Arc<GatewayState>) {
|
) -> (SocketAddr, Arc<GatewayState>) {
|
||||||
let state = Arc::new(GatewayState {
|
let state = Arc::new(GatewayState {
|
||||||
msg_tx: tokio::sync::RwLock::new(None),
|
msg_tx: tokio::sync::RwLock::new(None),
|
||||||
sse: Arc::new(SseManager::new()),
|
sse: SseManager::new(),
|
||||||
workspace: None,
|
workspace: None,
|
||||||
workspace_pool: None,
|
|
||||||
session_manager: None,
|
session_manager: None,
|
||||||
log_broadcaster: None,
|
log_broadcaster: None,
|
||||||
log_level_handle: None,
|
log_level_handle: None,
|
||||||
@@ -203,13 +202,13 @@ async fn start_test_server_with_provider(
|
|||||||
job_manager: None,
|
job_manager: None,
|
||||||
prompt_queue: None,
|
prompt_queue: None,
|
||||||
scheduler: None,
|
scheduler: None,
|
||||||
default_user_id: "test-user".to_string(),
|
user_id: "test-user".to_string(),
|
||||||
shutdown_tx: tokio::sync::RwLock::new(None),
|
shutdown_tx: tokio::sync::RwLock::new(None),
|
||||||
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
|
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
|
||||||
llm_provider: Some(llm_provider),
|
llm_provider: Some(llm_provider),
|
||||||
skill_registry: None,
|
skill_registry: None,
|
||||||
skill_catalog: None,
|
skill_catalog: None,
|
||||||
chat_rate_limiter: ironclaw::channels::web::server::PerUserRateLimiter::new(30, 60),
|
chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60),
|
||||||
oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
|
oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
|
||||||
webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
|
webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
|
||||||
registry_entries: Vec::new(),
|
registry_entries: Vec::new(),
|
||||||
@@ -219,12 +218,8 @@ async fn start_test_server_with_provider(
|
|||||||
active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(),
|
active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(),
|
||||||
});
|
});
|
||||||
|
|
||||||
let auth = ironclaw::channels::web::auth::MultiAuthState::single(
|
|
||||||
AUTH_TOKEN.to_string(),
|
|
||||||
"test-user".to_string(),
|
|
||||||
);
|
|
||||||
let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
|
let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
|
||||||
let bound_addr = start_server(addr, state.clone(), auth)
|
let bound_addr = start_server(addr, state.clone(), AUTH_TOKEN.to_string())
|
||||||
.await
|
.await
|
||||||
.expect("Failed to start test server");
|
.expect("Failed to start test server");
|
||||||
|
|
||||||
@@ -689,9 +684,8 @@ async fn test_no_llm_provider_returns_503() {
|
|||||||
// Create state WITHOUT llm_provider
|
// Create state WITHOUT llm_provider
|
||||||
let state = Arc::new(GatewayState {
|
let state = Arc::new(GatewayState {
|
||||||
msg_tx: tokio::sync::RwLock::new(None),
|
msg_tx: tokio::sync::RwLock::new(None),
|
||||||
sse: Arc::new(SseManager::new()),
|
sse: SseManager::new(),
|
||||||
workspace: None,
|
workspace: None,
|
||||||
workspace_pool: None,
|
|
||||||
session_manager: None,
|
session_manager: None,
|
||||||
log_broadcaster: None,
|
log_broadcaster: None,
|
||||||
log_level_handle: None,
|
log_level_handle: None,
|
||||||
@@ -701,13 +695,13 @@ async fn test_no_llm_provider_returns_503() {
|
|||||||
job_manager: None,
|
job_manager: None,
|
||||||
prompt_queue: None,
|
prompt_queue: None,
|
||||||
scheduler: None,
|
scheduler: None,
|
||||||
default_user_id: "test-user".to_string(),
|
user_id: "test-user".to_string(),
|
||||||
shutdown_tx: tokio::sync::RwLock::new(None),
|
shutdown_tx: tokio::sync::RwLock::new(None),
|
||||||
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
|
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
|
||||||
llm_provider: None, // No LLM!
|
llm_provider: None, // No LLM!
|
||||||
skill_registry: None,
|
skill_registry: None,
|
||||||
skill_catalog: None,
|
skill_catalog: None,
|
||||||
chat_rate_limiter: ironclaw::channels::web::server::PerUserRateLimiter::new(30, 60),
|
chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60),
|
||||||
oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
|
oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
|
||||||
webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
|
webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
|
||||||
registry_entries: Vec::new(),
|
registry_entries: Vec::new(),
|
||||||
@@ -717,12 +711,10 @@ async fn test_no_llm_provider_returns_503() {
|
|||||||
active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(),
|
active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(),
|
||||||
});
|
});
|
||||||
|
|
||||||
let auth = ironclaw::channels::web::auth::MultiAuthState::single(
|
|
||||||
AUTH_TOKEN.to_string(),
|
|
||||||
"test-user".to_string(),
|
|
||||||
);
|
|
||||||
let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
|
let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
|
||||||
let bound_addr = start_server(addr, state, auth).await.unwrap();
|
let bound_addr = start_server(addr, state, AUTH_TOKEN.to_string())
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
let url = format!("http://{}/v1/chat/completions", bound_addr);
|
let url = format!("http://{}/v1/chat/completions", bound_addr);
|
||||||
let resp = client()
|
let resp = client()
|
||||||
@@ -749,10 +741,9 @@ async fn test_chat_completions_body_too_large() {
|
|||||||
let state = ironclaw::channels::web::test_helpers::TestGatewayBuilder::new()
|
let state = ironclaw::channels::web::test_helpers::TestGatewayBuilder::new()
|
||||||
.llm_provider(llm_provider)
|
.llm_provider(llm_provider)
|
||||||
.build();
|
.build();
|
||||||
let auth_state = ironclaw::channels::web::auth::MultiAuthState::single(
|
let auth_state = ironclaw::channels::web::auth::AuthState {
|
||||||
AUTH_TOKEN.to_string(),
|
token: AUTH_TOKEN.to_string(),
|
||||||
"test-user".to_string(),
|
};
|
||||||
);
|
|
||||||
|
|
||||||
let app = Router::new()
|
let app = Router::new()
|
||||||
.route(
|
.route(
|
||||||
|
|||||||
@@ -13,11 +13,8 @@ use ironclaw::agent::routine_engine::RoutineEngine;
|
|||||||
use ironclaw::agent::{Agent, AgentDeps, SessionManager as AgentSessionManager};
|
use ironclaw::agent::{Agent, AgentDeps, SessionManager as AgentSessionManager};
|
||||||
use ironclaw::app::{AppBuilder, AppBuilderFlags};
|
use ironclaw::app::{AppBuilder, AppBuilderFlags};
|
||||||
use ironclaw::channels::IncomingMessage;
|
use ironclaw::channels::IncomingMessage;
|
||||||
use ironclaw::channels::web::auth::MultiAuthState;
|
|
||||||
use ironclaw::channels::web::log_layer::LogBroadcaster;
|
use ironclaw::channels::web::log_layer::LogBroadcaster;
|
||||||
use ironclaw::channels::web::server::{
|
use ironclaw::channels::web::server::{GatewayState, RateLimiter, start_server};
|
||||||
GatewayState, PerUserRateLimiter, RateLimiter, start_server,
|
|
||||||
};
|
|
||||||
use ironclaw::channels::web::sse::SseManager;
|
use ironclaw::channels::web::sse::SseManager;
|
||||||
use ironclaw::channels::web::ws::WsConnectionTracker;
|
use ironclaw::channels::web::ws::WsConnectionTracker;
|
||||||
use ironclaw::config::{Config, RegistryProviderConfig, RoutineConfig};
|
use ironclaw::config::{Config, RegistryProviderConfig, RoutineConfig};
|
||||||
@@ -214,9 +211,8 @@ impl GatewayWorkflowHarness {
|
|||||||
|
|
||||||
let gateway_state = Arc::new(GatewayState {
|
let gateway_state = Arc::new(GatewayState {
|
||||||
msg_tx: tokio::sync::RwLock::new(Some(gw_tx)),
|
msg_tx: tokio::sync::RwLock::new(Some(gw_tx)),
|
||||||
sse: Arc::new(SseManager::new()),
|
sse: SseManager::new(),
|
||||||
workspace: components.workspace.clone(),
|
workspace: components.workspace.clone(),
|
||||||
workspace_pool: None,
|
|
||||||
session_manager: Some(Arc::clone(&agent_session_manager)),
|
session_manager: Some(Arc::clone(&agent_session_manager)),
|
||||||
log_broadcaster: None,
|
log_broadcaster: None,
|
||||||
log_level_handle: None,
|
log_level_handle: None,
|
||||||
@@ -226,13 +222,13 @@ impl GatewayWorkflowHarness {
|
|||||||
job_manager: None,
|
job_manager: None,
|
||||||
prompt_queue: None,
|
prompt_queue: None,
|
||||||
scheduler: Some(scheduler_slot.clone()),
|
scheduler: Some(scheduler_slot.clone()),
|
||||||
default_user_id: user_id.clone(),
|
user_id: user_id.clone(),
|
||||||
shutdown_tx: tokio::sync::RwLock::new(None),
|
shutdown_tx: tokio::sync::RwLock::new(None),
|
||||||
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
|
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
|
||||||
llm_provider: Some(Arc::clone(&components.llm)),
|
llm_provider: Some(Arc::clone(&components.llm)),
|
||||||
skill_registry: components.skill_registry.clone(),
|
skill_registry: components.skill_registry.clone(),
|
||||||
skill_catalog: components.skill_catalog.clone(),
|
skill_catalog: components.skill_catalog.clone(),
|
||||||
chat_rate_limiter: PerUserRateLimiter::new(120, 60),
|
chat_rate_limiter: RateLimiter::new(120, 60),
|
||||||
oauth_rate_limiter: RateLimiter::new(10, 60),
|
oauth_rate_limiter: RateLimiter::new(10, 60),
|
||||||
webhook_rate_limiter: RateLimiter::new(10, 60),
|
webhook_rate_limiter: RateLimiter::new(10, 60),
|
||||||
registry_entries: Vec::new(),
|
registry_entries: Vec::new(),
|
||||||
@@ -258,7 +254,7 @@ impl GatewayWorkflowHarness {
|
|||||||
skills_config: components.config.skills.clone(),
|
skills_config: components.config.skills.clone(),
|
||||||
hooks: components.hooks,
|
hooks: components.hooks,
|
||||||
cost_guard: components.cost_guard,
|
cost_guard: components.cost_guard,
|
||||||
sse_tx: None,
|
sse_tx: Some(gateway_state.sse.sender()),
|
||||||
http_interceptor: None,
|
http_interceptor: None,
|
||||||
transcription: None,
|
transcription: None,
|
||||||
document_extraction: None,
|
document_extraction: None,
|
||||||
@@ -292,11 +288,10 @@ impl GatewayWorkflowHarness {
|
|||||||
}
|
}
|
||||||
|
|
||||||
let auth_token = "gateway-test-token".to_string();
|
let auth_token = "gateway-test-token".to_string();
|
||||||
let auth = MultiAuthState::single(auth_token.clone(), user_id.clone());
|
|
||||||
let addr = start_server(
|
let addr = start_server(
|
||||||
"127.0.0.1:0".parse().expect("valid localhost addr"),
|
"127.0.0.1:0".parse().expect("valid localhost addr"),
|
||||||
Arc::clone(&gateway_state),
|
Arc::clone(&gateway_state),
|
||||||
auth,
|
auth_token.clone(),
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
.expect("failed to start gateway server");
|
.expect("failed to start gateway server");
|
||||||
|
|||||||
@@ -39,9 +39,8 @@ async fn start_test_server() -> (
|
|||||||
|
|
||||||
let state = Arc::new(GatewayState {
|
let state = Arc::new(GatewayState {
|
||||||
msg_tx: tokio::sync::RwLock::new(Some(agent_tx)),
|
msg_tx: tokio::sync::RwLock::new(Some(agent_tx)),
|
||||||
sse: Arc::new(SseManager::new()),
|
sse: SseManager::new(),
|
||||||
workspace: None,
|
workspace: None,
|
||||||
workspace_pool: None,
|
|
||||||
session_manager: None,
|
session_manager: None,
|
||||||
log_broadcaster: None,
|
log_broadcaster: None,
|
||||||
log_level_handle: None,
|
log_level_handle: None,
|
||||||
@@ -51,13 +50,13 @@ async fn start_test_server() -> (
|
|||||||
job_manager: None,
|
job_manager: None,
|
||||||
prompt_queue: None,
|
prompt_queue: None,
|
||||||
scheduler: None,
|
scheduler: None,
|
||||||
default_user_id: "test-user".to_string(),
|
user_id: "test-user".to_string(),
|
||||||
shutdown_tx: tokio::sync::RwLock::new(None),
|
shutdown_tx: tokio::sync::RwLock::new(None),
|
||||||
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
|
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
|
||||||
llm_provider: None,
|
llm_provider: None,
|
||||||
skill_registry: None,
|
skill_registry: None,
|
||||||
skill_catalog: None,
|
skill_catalog: None,
|
||||||
chat_rate_limiter: ironclaw::channels::web::server::PerUserRateLimiter::new(30, 60),
|
chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60),
|
||||||
oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
|
oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
|
||||||
webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
|
webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
|
||||||
registry_entries: Vec::new(),
|
registry_entries: Vec::new(),
|
||||||
@@ -67,12 +66,8 @@ async fn start_test_server() -> (
|
|||||||
active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(),
|
active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(),
|
||||||
});
|
});
|
||||||
|
|
||||||
let auth = ironclaw::channels::web::auth::MultiAuthState::single(
|
|
||||||
AUTH_TOKEN.to_string(),
|
|
||||||
"test-user".to_string(),
|
|
||||||
);
|
|
||||||
let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
|
let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
|
||||||
let bound_addr = start_server(addr, state.clone(), auth)
|
let bound_addr = start_server(addr, state.clone(), AUTH_TOKEN.to_string())
|
||||||
.await
|
.await
|
||||||
.expect("Failed to start test server");
|
.expect("Failed to start test server");
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user