diff --git a/README.md b/README.md index c69cd71..bb597b6 100644 --- a/README.md +++ b/README.md @@ -264,6 +264,8 @@ Approved commands run through `bash -o errexit -o pipefail -c`. Stdout and stder Risk is model-generated guidance, not a security boundary. Read every proposed or edited command before approving it. +If `system_one_key`, `system_one_api`, and `system_one_model` are configured, CLAI uses that System One-compatible API for typed intent routing and command-risk auditing. The main LLM still generates commands and explanations, but System One decides whether a request is a command task, question, or history-clear request, and independently audits proposed command risk before `risk_appetite` can auto-run it. A higher System One risk label overrides the LLM risk label; low-confidence risk audits force a confirmation prompt. + ## Providers CLAI selects its native adapter from the configured `api` URL: @@ -310,6 +312,14 @@ json_mode=true reasoning=true ``` +Example System One configuration for TypeSafe Jev: + +```ini +system_one_key=ts-... +system_one_api=https://api.typesafe.ai/v1/systemone +system_one_model=jev-latest +``` + ## Configuration reference CLAI creates `~/.config/clai.cfg` on first use. It uses the established CLAI `key=value` format. The config path must be a regular file rather than a directory or symbolic link; CLAI enforces mode `0600` before reading it. @@ -326,6 +336,9 @@ CLAI creates `~/.config/clai.cfg` on first use. It uses the established CLAI `ke | `temp` | `0.1` | Sampling temperature. Invalid values fall back to `0.1`. | | `tokens` | `500` | Maximum requested output tokens. Invalid or non-positive values fall back to `500`. | | `reasoning` | empty | Optional reasoning-effort value; provider behavior is described above. | +| `system_one_key` | empty | Optional System One-compatible API credential for typed intent routing and risk auditing. | +| `system_one_api` | `https://api.typesafe.ai/v1/systemone` | System One-compatible HTTPS evaluation endpoint. Used only when key, API, and model are all present. Redirects must remain on the configured HTTPS origin (same host and port). | +| `system_one_model` | `jev-latest` | System One model used for typed judgments. | | `use_tools` | `false` | Opt in to discovering tools, sending their definitions to compatible providers, and allowing model-requested tool calls. | | `share_command_results` | `false` | Send bounded command results for immediate model interpretation and retain them for later context. | | `result_lines` | `20` | Maximum recent stdout and stderr lines stored for each shared result. | diff --git a/internal/app/app.go b/internal/app/app.go index f0fffe3..a647f1e 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -15,6 +15,7 @@ import ( "github.com/merefield/clai/internal/model" "github.com/merefield/clai/internal/provider" "github.com/merefield/clai/internal/runner" + "github.com/merefield/clai/internal/systemone" "github.com/merefield/clai/internal/ui" "github.com/merefield/clai/pkg/tool" ) @@ -24,6 +25,7 @@ type Application struct { History *history.Store Tools *tool.Registry Client provider.Client + SystemOne systemone.Client Runner runner.Runner UI *ui.Console ToolManager *mcptools.Manager @@ -53,7 +55,14 @@ func New(_ context.Context, in io.Reader, out, errOut io.Writer) (*Application, return nil, fmt.Errorf("initialize external tools: %w", err) } console := ui.New(in, out, errOut, cfg.HighContrast) - return &Application{Config: cfg, History: historyStore, Tools: toolManager.Registry(), Client: provider.New(cfg, nil), Runner: runner.Bash{Stdout: out, Stderr: errOut}, UI: console, ToolManager: toolManager}, nil + var systemOne systemone.Client + if strings.TrimSpace(cfg.SystemOneKey) != "" && strings.TrimSpace(cfg.SystemOneModel) != "" { + systemOne, err = systemone.New(cfg.SystemOneKey, cfg.SystemOneAPI, cfg.SystemOneModel, nil) + if err != nil { + return nil, err + } + } + return &Application{Config: cfg, History: historyStore, Tools: toolManager.Registry(), Client: provider.New(cfg, nil), SystemOne: systemOne, Runner: runner.Bash{Stdout: out, Stderr: errOut}, UI: console, ToolManager: toolManager}, nil } func (a *Application) Close() error { @@ -216,6 +225,14 @@ func (a *Application) setup() error { return err } a.Client = provider.New(a.Config, nil) + if strings.TrimSpace(a.Config.SystemOneKey) != "" && strings.TrimSpace(a.Config.SystemOneModel) != "" { + a.SystemOne, err = systemone.New(a.Config.SystemOneKey, a.Config.SystemOneAPI, a.Config.SystemOneModel, nil) + if err != nil { + return err + } + } else { + a.SystemOne = nil + } fmt.Fprintln(a.UI.Out, "CLAI configuration updated.") return nil } @@ -237,12 +254,12 @@ func (a *Application) process(ctx context.Context, query, requestedKind string) return err } } - kind := requestedKind - if kind == "" { - kind = "execute" - if isQuestion(query) { - kind = "question" - } + kind, err := a.routeIntent(ctx, query, requestedKind) + if err != nil { + return err + } + if kind == "clear_history" { + return a.clearHistory() } a.History.AppendText("user", query) previousResponseID := "" @@ -288,6 +305,10 @@ func (a *Application) process(ctx context.Context, query, requestedKind string) if HasPlaceholders(reply.Command) || HasPlaceholders(reply.Info) { reply = model.Reply{Info: "CLAI returned unresolved placeholders. Rephrase the request or specify missing values.", Risk: model.RiskNone, Variables: []model.Variable{}} } + forceConfirm, err := a.auditRisk(ctx, query, &reply) + if err != nil { + return err + } if err := a.History.AppendReply(reply); err != nil { return err } @@ -295,11 +316,84 @@ func (a *Application) process(ctx context.Context, query, requestedKind string) if reply.Command == "" { return nil } - return a.confirmAndRun(ctx, query, reply) + return a.confirmAndRun(ctx, query, reply, forceConfirm) } return fmt.Errorf("tool-call limit exceeded") } +func (a *Application) routeIntent(ctx context.Context, query, requestedKind string) (string, error) { + if requestedKind != "" { + return requestedKind, nil + } + if a.SystemOne != nil { + decision, err := a.SystemOne.RouteIntent(ctx, systemone.IntentRequest{UserRequest: query}) + if err != nil { + return "", fmt.Errorf("route intent with system one: %w", err) + } + switch decision.Intent { + case systemone.IntentQuestion: + return "question", nil + case systemone.IntentClearHistory: + if !systemone.ValidConfidence(decision.Confidence) || decision.Confidence < 0.65 { + return "", fmt.Errorf("system one history-clear intent is uncertain; use the explicit clear command to clear history") + } + return "clear_history", nil + case systemone.IntentExecute: + return "execute", nil + default: + return "", fmt.Errorf("system one returned unknown intent %q", decision.Intent) + } + } + if isQuestion(query) { + return "question", nil + } + return "execute", nil +} + +func (a *Application) auditRisk(ctx context.Context, query string, reply *model.Reply) (bool, error) { + if a.SystemOne == nil || reply.Command == "" { + return false, nil + } + decision, err := a.SystemOne.AuditRisk(ctx, systemone.RiskRequest{ + UserRequest: query, + Command: reply.Command, + Info: reply.Info, + LLMRisk: reply.Risk, + }) + if err != nil { + return false, fmt.Errorf("audit risk with system one: %w", err) + } + auditedRisk := NormalizeRisk(systemOneRisk(decision.Risk), reply.Command) + if riskRank(auditedRisk) > riskRank(reply.Risk) { + reply.Risk = auditedRisk + } + return !systemone.ValidConfidence(decision.Confidence) || decision.Confidence < 0.65, nil +} + +func systemOneRisk(value string) string { + switch value { + case "none": + return model.RiskNone + case "reversible_change": + return model.RiskReversible + case "danger_zone": + return model.RiskDanger + default: + return value + } +} + +func riskRank(value string) int { + switch value { + case model.RiskDanger: + return 2 + case model.RiskReversible: + return 1 + default: + return 0 + } +} + func (a *Application) reloadTools(ctx context.Context) error { if a.ToolManager == nil { return fmt.Errorf("external tool manager is unavailable") @@ -404,10 +498,10 @@ func (a *Application) resolveVariables(reply *model.Reply) error { return nil } -func (a *Application) confirmAndRun(ctx context.Context, originalQuery string, reply model.Reply) error { +func (a *Application) confirmAndRun(ctx context.Context, originalQuery string, reply model.Reply, forceConfirm bool) error { command := reply.Command edited := false - if RequiresConfirmation(reply.Risk, a.Config.RiskAppetite) { + if forceConfirm || RequiresConfirmation(reply.Risk, a.Config.RiskAppetite) { choice, err := a.UI.Choice("execute command? [y/e/N]: ") if err != nil { return err @@ -427,7 +521,7 @@ func (a *Application) confirmAndRun(ctx context.Context, originalQuery string, r a.UI.Cancel() return nil } - if reply.Risk == model.RiskDanger && a.Config.ConfirmDangerousCommands { + if (reply.Risk == model.RiskDanger || (forceConfirm && edited)) && a.Config.ConfirmDangerousCommands { confirm, err := a.UI.Choice("danger zone command, are you sure? [y/N]: ") if err != nil { return err diff --git a/internal/app/app_test.go b/internal/app/app_test.go index 7b23751..621e910 100644 --- a/internal/app/app_test.go +++ b/internal/app/app_test.go @@ -5,6 +5,8 @@ import ( "context" "encoding/json" "errors" + "fmt" + "math" "os" "os/exec" "path/filepath" @@ -16,6 +18,7 @@ import ( "github.com/merefield/clai/internal/mcptools" "github.com/merefield/clai/internal/model" "github.com/merefield/clai/internal/provider" + "github.com/merefield/clai/internal/systemone" "github.com/merefield/clai/internal/ui" "github.com/merefield/clai/pkg/tool" ) @@ -35,6 +38,23 @@ func (f *fakeClient) Complete(_ context.Context, request provider.Request) (prov func (f *fakeClient) SupportsTools() bool { return f.tools } +type fakeSystemOne struct { + intent systemone.IntentDecision + intentRequests []systemone.IntentRequest + risk systemone.RiskDecision + riskRequests []systemone.RiskRequest +} + +func (f *fakeSystemOne) RouteIntent(_ context.Context, request systemone.IntentRequest) (systemone.IntentDecision, error) { + f.intentRequests = append(f.intentRequests, request) + return f.intent, nil +} + +func (f *fakeSystemOne) AuditRisk(_ context.Context, request systemone.RiskRequest) (systemone.RiskDecision, error) { + f.riskRequests = append(f.riskRequests, request) + return f.risk, nil +} + type fakeRunner struct { results []model.CommandResult calls []string @@ -241,6 +261,136 @@ func TestProcessQuestionDoesNotRunCommand(t *testing.T) { } } +func TestSystemOneRoutesQuestionWithoutQuestionMark(t *testing.T) { + var out bytes.Buffer + client := &fakeClient{responses: []provider.Response{{Text: `{"cmd":"rm -rf /tmp/question-mode","info":"approximately 9.4248","risk":"danger zone","variables":[]}`, FinishReason: "stop"}}} + commandRunner := &fakeRunner{} + router := &fakeSystemOne{intent: systemone.IntentDecision{Intent: systemone.IntentQuestion, Confidence: 0.91}} + application := &Application{ + Config: &config.Config{Key: "test", Model: "test", API: "http://test", MaxHistoryTurns: 10}, + History: &history.Store{Path: filepath.Join(t.TempDir(), "history.json")}, + Tools: testTools(t), + Client: client, + SystemOne: router, + Runner: commandRunner, + UI: ui.New(strings.NewReader(""), &out, &out, false), + } + if err := application.process(context.Background(), "how much is 3 times pi", ""); err != nil { + t.Fatal(err) + } + if len(router.intentRequests) != 1 || router.intentRequests[0].UserRequest != "how much is 3 times pi" { + t.Fatalf("intent requests = %#v", router.intentRequests) + } + if len(commandRunner.calls) != 0 { + t.Fatalf("unexpected command: %#v", commandRunner.calls) + } + if !strings.Contains(out.String(), "approximately 9.4248") { + t.Fatalf("output = %q", out.String()) + } +} + +func TestSystemOneClearHistoryRequiresConfidentIntent(t *testing.T) { + for _, confidence := range []float64{0, 0.64, 0.65, 1, -1, 2, math.NaN(), math.Inf(1)} { + t.Run(fmt.Sprint(confidence), func(t *testing.T) { + var out bytes.Buffer + store := &history.Store{Path: filepath.Join(t.TempDir(), "history.json")} + store.AppendText("user", "keep this history") + if err := store.Save(10); err != nil { + t.Fatal(err) + } + before, err := os.ReadFile(store.Path) + if err != nil { + t.Fatal(err) + } + application := &Application{ + Config: &config.Config{}, History: store, Tools: testTools(t), + SystemOne: &fakeSystemOne{intent: systemone.IntentDecision{Intent: systemone.IntentClearHistory, Confidence: confidence}}, + UI: ui.New(strings.NewReader(""), &out, &out, false), + } + err = application.process(context.Background(), "forget that", "") + allowed := confidence == 0.65 || confidence == 1 + if allowed { + if err != nil || len(store.Messages) != 0 { + t.Fatalf("clear failed: %v", err) + } + } else { + if err == nil { + t.Fatal("uncertain intent accepted") + } + after, readErr := os.ReadFile(store.Path) + if readErr != nil || !bytes.Equal(before, after) || len(store.Messages) != 1 { + t.Fatalf("history changed on uncertain intent: %v", readErr) + } + } + }) + } +} + +func TestSystemOneRiskAuditUpgradesRiskAndPreventsAutoRun(t *testing.T) { + var out bytes.Buffer + client := &fakeClient{responses: []provider.Response{{Text: `{"cmd":"rm -rf /tmp/example","info":"removes files","risk":"none","variables":[]}`, FinishReason: "stop"}}} + commandRunner := &fakeRunner{} + auditor := &fakeSystemOne{ + intent: systemone.IntentDecision{Intent: systemone.IntentExecute, Confidence: 0.95}, + risk: systemone.RiskDecision{Risk: "danger_zone", Confidence: 0.96}, + } + application := &Application{ + Config: &config.Config{Key: "test", Model: "test", API: "http://test", RiskAppetite: 2, ConfirmDangerousCommands: true, MaxHistoryTurns: 10}, + History: &history.Store{Path: filepath.Join(t.TempDir(), "history.json")}, + Tools: testTools(t), + Client: client, + SystemOne: auditor, + Runner: commandRunner, + UI: ui.New(strings.NewReader("n\n"), &out, &out, true), + } + if err := application.process(context.Background(), "remove example", ""); err != nil { + t.Fatal(err) + } + if len(auditor.riskRequests) != 1 || auditor.riskRequests[0].LLMRisk != model.RiskNone { + t.Fatalf("risk requests = %#v", auditor.riskRequests) + } + if len(commandRunner.calls) != 0 { + t.Fatalf("dangerous command auto-ran: %#v", commandRunner.calls) + } + if !strings.Contains(out.String(), "DANGER ZONE") || !strings.Contains(out.String(), "[cancel]") { + t.Fatalf("output = %q", out.String()) + } +} + +func TestSystemOneLowConfidenceRiskForcesPrompt(t *testing.T) { + for _, confidence := range []float64{0.42, -0.1, 2, math.NaN(), math.Inf(1), math.Inf(-1)} { + for _, appetite := range []int{1, 2} { + t.Run(fmt.Sprintf("confidence=%v/appetite=%d", confidence, appetite), func(t *testing.T) { + var out bytes.Buffer + client := &fakeClient{responses: []provider.Response{{Text: `{"cmd":"printf ok","info":"prints ok","risk":"none","variables":[]}`, FinishReason: "stop"}}} + commandRunner := &fakeRunner{} + auditor := &fakeSystemOne{ + intent: systemone.IntentDecision{Intent: systemone.IntentExecute, Confidence: 0.95}, + risk: systemone.RiskDecision{Risk: "none", Confidence: confidence}, + } + application := &Application{ + Config: &config.Config{Key: "test", Model: "test", API: "http://test", RiskAppetite: appetite, MaxHistoryTurns: 10}, + History: &history.Store{Path: filepath.Join(t.TempDir(), "history.json")}, + Tools: testTools(t), + Client: client, + SystemOne: auditor, + Runner: commandRunner, + UI: ui.New(strings.NewReader("n\n"), &out, &out, true), + } + if err := application.process(context.Background(), "print ok", ""); err != nil { + t.Fatal(err) + } + if len(commandRunner.calls) != 0 { + t.Fatalf("low-confidence command auto-ran: %#v", commandRunner.calls) + } + if !strings.Contains(out.String(), "execute command?") || !strings.Contains(out.String(), "[cancel]") { + t.Fatalf("output = %q", out.String()) + } + }) + } + } +} + func TestEditedDangerousCommandStillRequiresDangerConfirmation(t *testing.T) { var out bytes.Buffer commandRunner := &fakeRunner{} @@ -251,7 +401,7 @@ func TestEditedDangerousCommandStillRequiresDangerConfirmation(t *testing.T) { } reply := model.Reply{Command: "rm -rf /tmp/example", Info: "removes the example", Risk: model.RiskDanger} - if err := application.confirmAndRun(context.Background(), "remove it", reply); err != nil { + if err := application.confirmAndRun(context.Background(), "remove it", reply, false); err != nil { t.Fatal(err) } if len(commandRunner.calls) != 0 { @@ -262,6 +412,33 @@ func TestEditedDangerousCommandStillRequiresDangerConfirmation(t *testing.T) { } } +func TestEditedForcedConfirmationRequiresDangerConfirmation(t *testing.T) { + for _, confirm := range []string{"n", "y"} { + t.Run(confirm, func(t *testing.T) { + var out bytes.Buffer + commandRunner := &fakeRunner{} + application := &Application{ + Config: &config.Config{RiskAppetite: 2, ConfirmDangerousCommands: true}, + Runner: commandRunner, + UI: ui.New(strings.NewReader("e\nrm -rf /tmp/example\n"+confirm+"\n"), &out, &out, true), + } + reply := model.Reply{Command: "printf ok", Risk: model.RiskNone} + if err := application.confirmAndRun(context.Background(), "print ok", reply, true); err != nil { + t.Fatal(err) + } + if !strings.Contains(out.String(), "danger zone command, are you sure?") { + t.Fatalf("missing danger confirmation: %s", out.String()) + } + if confirm == "n" && len(commandRunner.calls) != 0 { + t.Fatal("edited command ran after cancellation") + } + if confirm == "y" && (len(commandRunner.calls) != 1 || commandRunner.calls[0] != "rm -rf /tmp/example") { + t.Fatalf("runner calls = %v", commandRunner.calls) + } + }) + } +} + func TestShowHistorySanitizesStoredTerminalControls(t *testing.T) { var out bytes.Buffer application := &Application{ diff --git a/internal/config/config.go b/internal/config/config.go index 378692b..6190ecb 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -22,6 +22,9 @@ json_mode=false temp=0.1 tokens=500 reasoning= +system_one_key= +system_one_api=https://api.typesafe.ai/v1/systemone +system_one_model=jev-latest use_tools=false share_command_results=false result_lines=20 @@ -44,6 +47,9 @@ type Config struct { Temperature float64 Tokens int Reasoning string + SystemOneKey string + SystemOneAPI string + SystemOneModel string UseTools bool ShareCommandResults bool ResultLines int @@ -133,6 +139,9 @@ func (c *Config) refresh() { c.Temperature = floatValue(c.values["temp"], 0.1) c.Tokens = intValue(c.values["tokens"], 500, 1) c.Reasoning = c.values["reasoning"] + c.SystemOneKey = c.values["system_one_key"] + c.SystemOneAPI = stringValue(c.values["system_one_api"], "https://api.typesafe.ai/v1/systemone") + c.SystemOneModel = stringValue(c.values["system_one_model"], "jev-latest") c.UseTools = boolValue(c.values["use_tools"], false) c.ShareCommandResults = boolValue(c.values["share_command_results"], false) c.ResultLines = intValue(c.values["result_lines"], 20, 1) @@ -173,7 +182,7 @@ func (c *Config) Save() error { } func defaultKeys() []string { - return []string{"key", "hi_contrast", "expose_current_dir", "max_history_turns", "api", "model", "json_mode", "temp", "tokens", "reasoning", "use_tools", "share_command_results", "result_lines", "confirm_dangerous_commands", "risk_appetite", "exec_query", "question_query", "error_query"} + return []string{"key", "hi_contrast", "expose_current_dir", "max_history_turns", "api", "model", "json_mode", "temp", "tokens", "reasoning", "system_one_key", "system_one_api", "system_one_model", "use_tools", "share_command_results", "result_lines", "confirm_dangerous_commands", "risk_appetite", "exec_query", "question_query", "error_query"} } func atomicWrite(path string, data []byte, mode os.FileMode) error { diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 3529784..da95dca 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -23,6 +23,9 @@ func TestLoadCreatesCompatibleDefaults(t *testing.T) { if cfg.UseTools { t.Fatal("tools must be disabled by default") } + if cfg.SystemOneKey != "" || cfg.SystemOneAPI != "https://api.typesafe.ai/v1/systemone" || cfg.SystemOneModel != "jev-latest" { + t.Fatalf("unexpected system one defaults: %#v", cfg) + } info, err := os.Stat(path) if err != nil { t.Fatal(err) diff --git a/internal/systemone/client.go b/internal/systemone/client.go new file mode 100644 index 0000000..8a23742 --- /dev/null +++ b/internal/systemone/client.go @@ -0,0 +1,258 @@ +package systemone + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "math" + "net/http" + "net/url" + "strings" + "time" +) + +const ( + IntentExecute = "execute" + IntentQuestion = "question" + IntentClearHistory = "clear_history" +) + +type Client interface { + RouteIntent(context.Context, IntentRequest) (IntentDecision, error) + AuditRisk(context.Context, RiskRequest) (RiskDecision, error) +} + +type HTTPClient struct { + Key string + API string + Model string + Client *http.Client +} + +type IntentRequest struct { + UserRequest string `json:"user_request"` +} + +type IntentDecision struct { + Intent string + Confidence float64 +} + +type RiskRequest struct { + UserRequest string `json:"user_request"` + Command string `json:"command"` + Info string `json:"info"` + LLMRisk string `json:"llm_risk"` +} + +type RiskDecision struct { + Risk string + Confidence float64 +} + +type question map[string]any + +type request struct { + State any `json:"state"` + Model string `json:"model"` + Questions map[string]question `json:"questions"` +} + +type response struct { + Answers map[string]answer `json:"answers"` +} + +type answer struct { + Type string `json:"type"` + Choice string `json:"choice,omitempty"` + Confidence *float64 `json:"confidence,omitempty"` + Prob map[string]*float64 `json:"probabilities,omitempty"` + Noul float64 `json:"noul,omitempty"` +} + +func Configured(key, api, model string) bool { + return strings.TrimSpace(key) != "" && validEndpoint(api) && strings.TrimSpace(model) != "" +} + +func validEndpoint(api string) bool { + endpoint, err := url.Parse(api) + return err == nil && endpoint.Scheme == "https" && endpoint.Hostname() != "" && endpoint.User == nil && endpoint.Fragment == "" +} + +func ValidConfidence(value float64) bool { + return !math.IsNaN(value) && value >= 0 && value <= 1 +} + +func New(key, api, model string, client *http.Client) (*HTTPClient, error) { + if !validEndpoint(api) { + return nil, fmt.Errorf("system_one_api must be a valid HTTPS endpoint without user info or fragment") + } + if client == nil { + client = &http.Client{Timeout: 30 * time.Second} + } + secureClient := *client + origin, _ := url.Parse(api) + secureClient.CheckRedirect = func(req *http.Request, via []*http.Request) error { + if !validEndpoint(req.URL.String()) { + return fmt.Errorf("system one redirect requires HTTPS") + } + if !strings.EqualFold(req.URL.Hostname(), origin.Hostname()) || httpsPort(req.URL) != httpsPort(origin) { + return fmt.Errorf("system one redirect must remain on the configured origin") + } + if client.CheckRedirect != nil { + return client.CheckRedirect(req, via) + } + if len(via) >= 10 { + return fmt.Errorf("stopped after 10 redirects") + } + return nil + } + return &HTTPClient{Key: key, API: api, Model: model, Client: &secureClient}, nil +} + +func httpsPort(endpoint *url.URL) string { + if port := endpoint.Port(); port != "" { + return port + } + return "443" +} + +func (c *HTTPClient) RouteIntent(ctx context.Context, input IntentRequest) (IntentDecision, error) { + payload := request{ + State: input, + Model: c.Model, + Questions: map[string]question{ + "intent": { + "type": "choice", + "instructions": "Which CLAI workflow should handle this user request?", + "criteria": map[string]string{ + IntentExecute: "The user wants CLAI to propose a shell command or perform a terminal task.", + IntentQuestion: "The user asks for an explanation or answer and no shell command should be proposed.", + IntentClearHistory: "The user wants to clear, reset, forget, or flush CLAI conversation history.", + }, + }, + }, + } + answers, err := c.evaluate(ctx, payload) + if err != nil { + return IntentDecision{}, err + } + value, ok := answers["intent"] + if !ok || value.Type != "choice" || value.Choice == "" { + return IntentDecision{}, fmt.Errorf("system one intent response missing choice") + } + return IntentDecision{Intent: value.Choice, Confidence: *value.Confidence}, nil +} + +func (c *HTTPClient) AuditRisk(ctx context.Context, input RiskRequest) (RiskDecision, error) { + payload := request{ + State: input, + Model: c.Model, + Questions: map[string]question{ + "risk": { + "type": "choice", + "instructions": "Classify the risk of running the proposed shell command for this user request.", + "criteria": map[string]string{ + "none": "Read-only inspection or harmless output; no meaningful filesystem, process, network, or configuration change.", + "reversible_change": "A change that is normally undoable, such as creating files, installing dependencies, editing working-tree files, or switching branches.", + "danger_zone": "Deletion, overwrite, reset, force, permission, credential, destructive, privileged, or hard-to-reverse changes.", + }, + }, + }, + } + answers, err := c.evaluate(ctx, payload) + if err != nil { + return RiskDecision{}, err + } + value, ok := answers["risk"] + if !ok || value.Type != "choice" || value.Choice == "" { + return RiskDecision{}, fmt.Errorf("system one risk response missing choice") + } + return RiskDecision{Risk: value.Choice, Confidence: *value.Confidence}, nil +} + +func (c *HTTPClient) evaluate(ctx context.Context, payload request) (map[string]answer, error) { + if !validEndpoint(c.API) { + return nil, fmt.Errorf("system_one_api must be a valid HTTPS endpoint") + } + body, err := json.Marshal(payload) + if err != nil { + return nil, fmt.Errorf("encode system one request: %w", err) + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.API, bytes.NewReader(body)) + if err != nil { + return nil, err + } + req.Header.Set("Authorization", "Bearer "+c.Key) + req.Header.Set("Content-Type", "application/json") + resp, err := c.Client.Do(req) + if err != nil { + return nil, fmt.Errorf("system one request failed: %w", err) + } + defer resp.Body.Close() + responseBody, err := io.ReadAll(io.LimitReader(resp.Body, 4<<20)) + if err != nil { + return nil, fmt.Errorf("read system one response: %w", err) + } + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + return nil, fmt.Errorf("system one request failed (HTTP %d): %q", resp.StatusCode, strings.TrimSpace(string(responseBody))) + } + var decoded response + if err := json.Unmarshal(responseBody, &decoded); err != nil { + return nil, fmt.Errorf("parse system one response: %w", err) + } + if len(decoded.Answers) == 0 { + return nil, fmt.Errorf("system one response returned no answers") + } + for name, question := range payload.Questions { + value, present := decoded.Answers[name] + criteria, ok := question["criteria"].(map[string]string) + if !present || !ok || value.Type != "choice" { + return nil, fmt.Errorf("system one %s response missing valid choice", name) + } + if err := validateChoice(value, criteria); err != nil { + return nil, fmt.Errorf("system one %s response: %w", name, err) + } + } + return decoded.Answers, nil +} + +func validateChoice(value answer, criteria map[string]string) error { + if _, ok := criteria[value.Choice]; !ok { + return fmt.Errorf("choice is not a submitted criterion") + } + if value.Confidence == nil || !ValidConfidence(*value.Confidence) { + return fmt.Errorf("confidence must be present and within [0, 1]") + } + if len(value.Prob) != len(criteria) { + return fmt.Errorf("probabilities must include exactly the submitted criteria") + } + total := 0.0 + twoDecimal := true + for option := range criteria { + probability := value.Prob[option] + if probability == nil || !ValidConfidence(*probability) { + return fmt.Errorf("probability for %q must be present and within [0, 1]", option) + } + total += *probability + if math.Abs(*probability*100-math.Round(*probability*100)) > 1e-9 { + twoDecimal = false + } + } + // Two-decimal distributions can lose up to half a hundredth per option. + tolerance := 1e-6 + if twoDecimal { + tolerance += 0.005 * float64(len(criteria)) + } + if math.Abs(total-1) > tolerance { + return fmt.Errorf("probabilities must sum to 1") + } + for _, probability := range value.Prob { + if *probability > *value.Prob[value.Choice] { + return fmt.Errorf("choice must have the highest probability") + } + } + return nil +} diff --git a/internal/systemone/client_test.go b/internal/systemone/client_test.go new file mode 100644 index 0000000..b22ad6a --- /dev/null +++ b/internal/systemone/client_test.go @@ -0,0 +1,312 @@ +package systemone + +import ( + "context" + "encoding/json" + "fmt" + "io" + "math" + "net/http" + "strings" + "testing" +) + +type roundTripFunc func(*http.Request) (*http.Response, error) + +func (f roundTripFunc) RoundTrip(request *http.Request) (*http.Response, error) { return f(request) } + +func TestConfiguredRequiresKeyAPIAndModel(t *testing.T) { + if Configured("", "https://api.typesafe.ai/v1/systemone", "jev-latest") { + t.Fatal("empty key should not be configured") + } + if !Configured("key", "https://api.typesafe.ai/v1/systemone", "jev-latest") { + t.Fatal("complete system one config should be configured") + } +} + +func TestRejectsInsecureEndpoints(t *testing.T) { + for _, endpoint := range []string{"http://example.test/v1/systemone", "http://localhost:8080", "", "example.test", "/v1/systemone", "https:///missing-host", "https://user:pass@example.test", "https://example.test/#fragment", "https://example.test:bad"} { + t.Run(endpoint, func(t *testing.T) { + if Configured("key", endpoint, "model") { + t.Fatal("insecure endpoint enabled") + } + if _, err := New("key", endpoint, "model", nil); err == nil { + t.Fatal("constructor accepted insecure endpoint") + } + client := &HTTPClient{Key: "key", API: endpoint} + if _, err := client.AuditRisk(context.Background(), RiskRequest{}); err == nil { + t.Fatal("request accepted insecure endpoint") + } + }) + } +} + +func TestRedirectRequiresHTTPS(t *testing.T) { + for _, scheme := range []string{"http", "https"} { + t.Run(scheme, func(t *testing.T) { + calls := 0 + httpClient := &http.Client{Transport: roundTripFunc(func(r *http.Request) (*http.Response, error) { + calls++ + if calls == 1 { + return &http.Response{StatusCode: http.StatusFound, Header: http.Header{"Location": []string{scheme + "://example.test/next"}}, Body: io.NopCloser(strings.NewReader(""))}, nil + } + return jsonResponse(`{"answers":{"risk":{"type":"choice","choice":"none","confidence":1,"probabilities":{"none":1,"reversible_change":0,"danger_zone":0}}}}`), nil + })} + client, err := New("secret", "https://example.test/start", "model", httpClient) + if err != nil { + t.Fatal(err) + } + _, err = client.AuditRisk(context.Background(), RiskRequest{}) + if scheme == "http" && (err == nil || calls != 1) { + t.Fatalf("plaintext redirect: calls=%d error=%v", calls, err) + } + if scheme == "https" && (err != nil || calls != 2) { + t.Fatalf("HTTPS redirect: calls=%d error=%v", calls, err) + } + if httpClient.CheckRedirect != nil { + t.Fatal("modified caller's HTTP client") + } + }) + } +} + +func TestRedirectRequiresConfiguredOrigin(t *testing.T) { + for _, tc := range []struct { + location string + allowed bool + }{ + {"/next", true}, + {"https://example.test/next", true}, + {"https://EXAMPLE.test:443/next", true}, + {"https://other.test/next", false}, + {"https://sub.example.test/next", false}, + {"https://example.test:8443/next", false}, + } { + t.Run(tc.location, func(t *testing.T) { + calls := 0 + policyCalls := 0 + httpClient := &http.Client{ + CheckRedirect: func(*http.Request, []*http.Request) error { policyCalls++; return nil }, + Transport: roundTripFunc(func(r *http.Request) (*http.Response, error) { + calls++ + if calls == 1 { + return &http.Response{StatusCode: http.StatusFound, Header: http.Header{"Location": []string{tc.location}}, Body: io.NopCloser(strings.NewReader(""))}, nil + } + return jsonResponse(`{"answers":{"risk":{"type":"choice","choice":"none","confidence":1,"probabilities":{"none":1,"reversible_change":0,"danger_zone":0}}}}`), nil + }), + } + client, err := New("secret", "https://example.test/start", "model", httpClient) + if err != nil { + t.Fatal(err) + } + _, err = client.AuditRisk(context.Background(), RiskRequest{}) + if tc.allowed { + if err != nil || calls != 2 || policyCalls != 1 { + t.Fatalf("same-origin redirect: calls=%d policy=%d err=%v", calls, policyCalls, err) + } + } else if err == nil || calls != 1 || policyCalls != 0 { + t.Fatalf("cross-origin redirect was not blocked before sending: calls=%d policy=%d err=%v", calls, policyCalls, err) + } + }) + } +} + +func TestRouteIntentUsesChoiceQuestion(t *testing.T) { + var request map[string]any + httpClient := &http.Client{Transport: roundTripFunc(func(r *http.Request) (*http.Response, error) { + if r.Header.Get("Authorization") != "Bearer system-key" { + t.Errorf("authorization = %q", r.Header.Get("Authorization")) + } + if err := json.NewDecoder(r.Body).Decode(&request); err != nil { + t.Error(err) + } + body := `{"answers":{"intent":{"type":"choice","choice":"question","confidence":0.88,"probabilities":{"question":0.88,"execute":0.11,"clear_history":0.01}}}}` + return jsonResponse(body), nil + })} + client, err := New("system-key", "https://system-one.example/v1/systemone", "jev-test", httpClient) + if err != nil { + t.Fatal(err) + } + decision, err := client.RouteIntent(context.Background(), IntentRequest{UserRequest: "how much is 3*pi"}) + if err != nil { + t.Fatal(err) + } + if decision.Intent != IntentQuestion || decision.Confidence != 0.88 { + t.Fatalf("decision = %#v", decision) + } + if request["model"] != "jev-test" { + t.Fatalf("request = %#v", request) + } + questions := request["questions"].(map[string]any) + intent := questions["intent"].(map[string]any) + if intent["type"] != "choice" { + t.Fatalf("intent question = %#v", intent) + } +} + +func TestAuditRiskUsesChoiceQuestion(t *testing.T) { + var request map[string]any + httpClient := &http.Client{Transport: roundTripFunc(func(r *http.Request) (*http.Response, error) { + if err := json.NewDecoder(r.Body).Decode(&request); err != nil { + t.Error(err) + } + body := `{"answers":{"risk":{"type":"choice","choice":"danger_zone","confidence":0.93,"probabilities":{"danger_zone":0.93,"reversible_change":0.06,"none":0.01}}}}` + return jsonResponse(body), nil + })} + client, err := New("system-key", "https://system-one.example/v1/systemone", "jev-test", httpClient) + if err != nil { + t.Fatal(err) + } + decision, err := client.AuditRisk(context.Background(), RiskRequest{UserRequest: "remove it", Command: "rm -rf tmp", LLMRisk: "none"}) + if err != nil { + t.Fatal(err) + } + if decision.Risk != "danger_zone" || decision.Confidence != 0.93 { + t.Fatalf("decision = %#v", decision) + } + state := request["state"].(map[string]any) + if state["command"] != "rm -rf tmp" || state["llm_risk"] != "none" { + t.Fatalf("state = %#v", state) + } +} + +func jsonResponse(body string) *http.Response { + return &http.Response{StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"application/json"}}, Body: io.NopCloser(strings.NewReader(body))} +} + +func TestConfidenceValidation(t *testing.T) { + for _, value := range []float64{-1, 2, math.NaN(), math.Inf(1), math.Inf(-1)} { + if ValidConfidence(value) { + t.Errorf("accepted invalid confidence %v", value) + } + } + for _, value := range []float64{0, 0.65, 1} { + if !ValidConfidence(value) { + t.Errorf("rejected valid confidence %v", value) + } + } +} + +func TestResponseConfidenceRange(t *testing.T) { + for _, kind := range []string{"intent", "risk"} { + for _, confidence := range []string{"-0.1", "2", "0", "1"} { + t.Run(kind+"/"+confidence, func(t *testing.T) { + choice, probabilities := "none", `{"none":1,"reversible_change":0,"danger_zone":0}` + if kind == "intent" { + choice, probabilities = "execute", `{"execute":1,"question":0,"clear_history":0}` + } + httpClient := &http.Client{Transport: roundTripFunc(func(r *http.Request) (*http.Response, error) { + return jsonResponse(fmt.Sprintf(`{"answers":{"%s":{"type":"choice","choice":%q,"confidence":%s,"probabilities":%s}}}`, kind, choice, confidence, probabilities)), nil + })} + client, err := New("key", "https://example.test", "model", httpClient) + if err != nil { + t.Fatal(err) + } + if kind == "intent" { + _, err = client.RouteIntent(context.Background(), IntentRequest{}) + } else { + _, err = client.AuditRisk(context.Background(), RiskRequest{}) + } + invalid := confidence == "-0.1" || confidence == "2" + if (err != nil) != invalid { + t.Fatalf("confidence %s: error = %v", confidence, err) + } + }) + } + } +} + +func TestChoiceResponseValidation(t *testing.T) { + for _, kind := range []string{"risk", "intent"} { + for _, tc := range []struct { + name string + choice string + confidence string + probabilities string + valid bool + }{ + {"missing distribution", "selected", "1", "", false}, + {"null distribution", "selected", "1", "null", false}, + {"empty distribution", "selected", "1", `{}`, false}, + {"missing option", "selected", "1", `{"selected":1,"other":0}`, false}, + {"unknown option", "selected", "1", `{"selected":1,"other":0,"unknown":0}`, false}, + {"extra option", "selected", "1", `{"selected":1,"other":0,"last":0,"unknown":0}`, false}, + {"null probability", "selected", "1", `{"selected":1,"other":null,"last":0}`, false}, + {"negative probability", "selected", "1", `{"selected":1,"other":-0.1,"last":0.1}`, false}, + {"large probability", "selected", "1", `{"selected":2,"other":0,"last":0}`, false}, + {"invalid sum", "selected", "1", `{"selected":0.5,"other":0,"last":0}`, false}, + {"unknown choice", "unknown", "1", `{"selected":1,"other":0,"last":0}`, false}, + {"nonwinning choice", "other", "1", `{"selected":1,"other":0,"last":0}`, false}, + {"null confidence", "selected", "null", `{"selected":1,"other":0,"last":0}`, false}, + {"missing confidence", "selected", "", `{"selected":1,"other":0,"last":0}`, false}, + {"valid", "selected", "1", `{"selected":1,"other":0,"last":0}`, true}, + {"tie", "selected", "0", `{"selected":0.5,"other":0.5,"last":0}`, true}, + {"rounding", "selected", "0", `{"selected":0.3333333,"other":0.3333333,"last":0.3333333}`, true}, + {"two decimal sum below one", "selected", "0", `{"selected":0.33,"other":0.33,"last":0.33}`, true}, + {"two decimal sum above one", "selected", "0", `{"selected":0.34,"other":0.34,"last":0.33}`, true}, + {"two decimal sum too low", "selected", "0", `{"selected":0.34,"other":0.32,"last":0.32}`, false}, + {"two decimal sum too high", "selected", "0", `{"selected":0.34,"other":0.34,"last":0.34}`, false}, + {"higher precision stays strict", "selected", "0", `{"selected":0.333,"other":0.333,"last":0.333}`, false}, + } { + t.Run(kind+"/"+tc.name, func(t *testing.T) { + replacer := strings.NewReplacer("selected", "none", "other", "reversible_change", "last", "danger_zone") + if kind == "intent" { + replacer = strings.NewReplacer("selected", "execute", "other", "question", "last", "clear_history") + } + body := fmt.Sprintf(`{"answers":{"%s":{"type":"choice","choice":%q`, kind, replacer.Replace(tc.choice)) + if tc.confidence != "" { + body += `,"confidence":` + tc.confidence + } + if tc.probabilities != "" { + body += `,"probabilities":` + replacer.Replace(tc.probabilities) + } + body += `}}}` + client, err := New("key", "https://example.test", "model", &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + return jsonResponse(body), nil + })}) + if err != nil { + t.Fatal(err) + } + if kind == "risk" { + _, err = client.AuditRisk(context.Background(), RiskRequest{}) + } else { + _, err = client.RouteIntent(context.Background(), IntentRequest{}) + } + if (err == nil) != tc.valid { + t.Fatalf("valid=%v error=%v", tc.valid, err) + } + }) + } + } +} + +func TestHTTPErrorEscapesTerminalControls(t *testing.T) { + for _, body := range []string{ + "service unavailable", + "failure\x1b[2J\x1b[Hforged success", + "failure\x1b]52;c;c2VjcmV0\a", + "failure\rforged\nmessage\b\t\x00", + "failure\u009b2J\u009d52;c;data\u009c", + } { + t.Run(fmt.Sprintf("%q", body), func(t *testing.T) { + client, err := New("key", "https://example.test", "model", &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + response := jsonResponse(body) + response.StatusCode = http.StatusBadGateway + return response, nil + })}) + if err != nil { + t.Fatal(err) + } + _, err = client.AuditRisk(context.Background(), RiskRequest{}) + want := fmt.Sprintf("system one request failed (HTTP 502): %q", strings.TrimSpace(body)) + if err == nil || err.Error() != want { + t.Fatalf("error = %v; want %s", err, want) + } + for _, r := range err.Error() { + if r < 32 || (r >= 127 && r <= 159) { + t.Fatalf("raw terminal control %U in error", r) + } + } + }) + } +}