Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
44 changes: 42 additions & 2 deletions apps/desktop/src-tauri/src/sidecar.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3061,6 +3061,8 @@ fn validate_privacy_policy_value(policy: &serde_json::Value) -> Result<(), Strin
"response_restore",
"restore_tool_arguments",
"placeholder_notice",
"skip_tool_declarations",
"inspect_additional_tools",
],
&[
"id",
Expand Down Expand Up @@ -3130,7 +3132,12 @@ fn validate_privacy_policy_value(policy: &serde_json::Value) -> Result<(), Strin
if let Some(allowlist_rules) = object.get("allowlist_rules") {
validate_privacy_allowlist_rules(allowlist_rules)?;
}
for field in ["restore_tool_arguments", "placeholder_notice"] {
for field in [
"restore_tool_arguments",
"placeholder_notice",
"skip_tool_declarations",
"inspect_additional_tools",
] {
if let Some(value) = object.get(field) {
if !value.is_boolean() {
return Err(format!("privacy policy {field} must be a boolean"));
Expand Down Expand Up @@ -3248,6 +3255,11 @@ fn validate_privacy_policy_patch(patch: serde_json::Value) -> Result<serde_json:
"response_restore" if value.is_boolean() => {}
"restore_tool_arguments" if value.is_boolean() => {}
"placeholder_notice" if value.is_boolean() => {}
"skip_tool_declarations" | "inspect_additional_tools" => {
if !value.is_boolean() {
return Err(format!("privacy policy {field} patch must be a boolean"));
}
}
"enabled" => {
return Err("privacy policy enabled patch must be a boolean".to_string());
}
Expand Down Expand Up @@ -5906,10 +5918,38 @@ mod tests {
"response_action": "allow",
"response_restore": true,
"restore_tool_arguments": true,
"placeholder_notice": true
"placeholder_notice": true,
"skip_tool_declarations": false,
"inspect_additional_tools": false
})
}

#[test]
fn privacy_policy_tool_declarations_require_booleans() {
for field in ["skip_tool_declarations", "inspect_additional_tools"] {
for value in [serde_json::json!(false), serde_json::json!(true)] {
let mut policy = privacy_policy_value();
policy[field] = value.clone();
assert_eq!(
parse_privacy_policy(&serde_json::to_vec(&policy).unwrap()).unwrap()[field],
value
);
let patch = serde_json::json!({field: value});
assert_eq!(validate_privacy_policy_patch(patch.clone()).unwrap(), patch);
}
for value in [
serde_json::Value::Null,
serde_json::json!("false"),
serde_json::json!(0),
] {
let mut policy = privacy_policy_value();
policy[field] = value.clone();
assert!(parse_privacy_policy(&serde_json::to_vec(&policy).unwrap()).is_err());
assert!(validate_privacy_policy_patch(serde_json::json!({field: value})).is_err());
}
}
}

#[test]
fn strictly_parses_the_privacy_policy_control_contract() {
let page = serde_json::to_vec(&serde_json::json!({
Expand Down
11 changes: 11 additions & 0 deletions core/internal/storage/migrate/defaults.go
Original file line number Diff line number Diff line change
Expand Up @@ -641,5 +641,16 @@ SET document_json = json_remove(document_json, '$.disabled_models')`,
`ALTER TABLE billing_ledger ADD COLUMN local_access_token_id TEXT`,
`CREATE INDEX request_records_root_token_time_idx ON request_records(local_access_token_id, started_at DESC, id DESC) WHERE parent_request_id IS NULL`,
}},
{Version: 34, Name: "privacy_tool_declaration_defaults", Statements: []string{
`UPDATE policies
SET document_json = json_insert(
document_json,
'$.skip_tool_declarations', json('false'),
'$.inspect_additional_tools', json('false')
)
WHERE id = 'policy_privacy_default'
AND (json_type(document_json, '$.skip_tool_declarations') IS NULL
OR json_type(document_json, '$.inspect_additional_tools') IS NULL)`,
}},
}
}
83 changes: 83 additions & 0 deletions core/internal/storage/sqlite/policy_store_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package sqlite

import (
"context"
"database/sql"
"encoding/json"
"errors"
"path/filepath"
Expand All @@ -10,8 +11,90 @@ import (

"github.com/QuantumNous/astrlink/core/contract"
storagecontract "github.com/QuantumNous/astrlink/core/internal/storage"
"github.com/QuantumNous/astrlink/core/internal/storage/migrate"
)

func TestPrivacyToolDeclarationMigration(t *testing.T) {
for _, fields := range []string{
`{}`,
`{"skip_tool_declarations":true}`,
`{"inspect_additional_tools":true}`,
`{"skip_tool_declarations":false,"inspect_additional_tools":false}`,
`{"skip_tool_declarations":true,"inspect_additional_tools":true}`,
} {
t.Run(fields, func(t *testing.T) {
ctx := context.Background()
path := filepath.Join(t.TempDir(), "astrlink.db")
database, err := sql.Open(driverName, path)
if err != nil {
t.Fatal(err)
}
defer database.Close()
runner, err := migrate.New(migrate.SQLDatabase{DB: database}, migrate.DefaultMigrations()[:33])
if err != nil {
t.Fatal(err)
}
if err := runner.Up(ctx); err != nil {
t.Fatal(err)
}
if _, err := database.Exec(`UPDATE policies SET document_json = json_patch(
json_set(document_json, '$.enabled', json('true'), '$.min_confidence', 0.85), ?)
WHERE id = 'policy_privacy_default'`, fields); err != nil {
t.Fatal(err)
}
var original string
if err := database.QueryRow(`SELECT document_json FROM policies WHERE id = 'policy_privacy_default'`).Scan(&original); err != nil {
t.Fatal(err)
}
var want map[string]json.RawMessage
if err := json.Unmarshal([]byte(original), &want); err != nil {
t.Fatal(err)
}
for _, field := range []string{"skip_tool_declarations", "inspect_additional_tools"} {
if _, exists := want[field]; !exists {
want[field] = json.RawMessage(`false`)
}
}
if err := database.Close(); err != nil {
t.Fatal(err)
}

store := openTestStore(t, path)
defer store.Close()
var migrated string
if err := store.db.QueryRow(`SELECT document_json FROM policies WHERE id = 'policy_privacy_default'`).Scan(&migrated); err != nil {
t.Fatal(err)
}
var got map[string]json.RawMessage
if err := json.Unmarshal([]byte(migrated), &got); err != nil {
t.Fatal(err)
}
if !reflect.DeepEqual(got, want) {
t.Fatalf("migrated policy = %s, want original settings with missing tool fields defaulted to false", migrated)
}
record, err := store.GetPolicy(ctx, contract.DefaultPrivacyPolicyID)
if err != nil {
t.Fatal(err)
}
record.Policy.SkipToolDeclarations = !record.Policy.SkipToolDeclarations
record.Policy.InspectAdditionalTools = !record.Policy.InspectAdditionalTools
updated, err := store.UpdatePolicy(ctx, record.Policy, record.ETag)
if err != nil {
t.Fatalf("save migrated policy: %v", err)
}
if err := store.Close(); err != nil {
t.Fatal(err)
}
store = openTestStore(t, path)
defer store.Close()
reloaded, err := store.GetPolicy(ctx, contract.DefaultPrivacyPolicyID)
if err != nil || !reflect.DeepEqual(reloaded, updated) {
t.Fatalf("policy after restart = %#v, %v, want %#v", reloaded, err, updated)
}
})
}
}

func TestDefaultPrivacyPolicyMigrationAndETagUpdate(t *testing.T) {
databasePath := filepath.Join(t.TempDir(), "astrlink.db")
store := openTestStore(t, databasePath)
Expand Down
Loading