mirror of
https://github.com/Ed1s0nZ/CyberStrikeAI.git
synced 2026-10-01 21:50:20 +02:00
feat: add TypeSafe Jev as HITL audit backend
Let audit_agent approve or reject with one System One call instead of chat JSON, while keeping the OpenAI-compatible backend as an option. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
1 parent
3aa9274675
commit
38b96ec67a
30 files changed
+1372
-21
No files matched your search
@@ -1073,6 +1073,7 @@ func setupRoutes(
|
||||
protected.PUT("/config", configHandler.UpdateConfig)
|
||||
protected.POST("/config/apply", configHandler.ApplyConfig)
|
||||
protected.POST("/config/test-openai", configHandler.TestOpenAI)
|
||||
protected.POST("/config/test-typesafe", configHandler.TestTypeSafe)
|
||||
protected.POST("/config/test-vision", configHandler.TestVision)
|
||||
protected.POST("/config/list-models", configHandler.ListModels)
|
||||
|
||||
|
||||
@@ -1114,7 +1114,9 @@ type AgentConfig struct {
|
||||
// tool_whitelist 可在侧栏「应用」时合并写入 config.yaml 并立即生效。
|
||||
// audit_agent_prompt / audit_agent_prompt_review_edit 可在人机协同页编辑并立即生效;空则使用内置默认。
|
||||
type HitlConfig struct {
|
||||
// AuditModel 审计 Agent 专用模型;字段留空时继承 OpenAI 主配置,便于用小模型做审批。
|
||||
// AuditBackend 审计 Agent 后端:openai(兼容协议聊天模型)或 typesafe(Jev 结构化裁决)。空值视为 openai。
|
||||
AuditBackend string `yaml:"audit_backend,omitempty" json:"audit_backend,omitempty"`
|
||||
// AuditModel 审计 Agent 专用模型。openai 后端空字段继承主模型;typesafe 后端 api_key 必填,不继承主模型密钥。
|
||||
AuditModel OpenAIConfig `yaml:"audit_model,omitempty" json:"audit_model,omitempty"`
|
||||
// ToolWhitelist 全局免审批工具名(与白名单内工具不触发 HITL 审批)。
|
||||
ToolWhitelist []string `yaml:"tool_whitelist,omitempty" json:"tool_whitelist,omitempty"`
|
||||
@@ -1176,6 +1178,37 @@ func (h HitlConfig) RetentionDaysEffective() int {
|
||||
return *h.RetentionDays
|
||||
}
|
||||
|
||||
const (
|
||||
HitlAuditBackendOpenAI = "openai"
|
||||
HitlAuditBackendTypeSafe = "typesafe"
|
||||
TypeSafeDefaultBaseURL = "https://api.typesafe.ai"
|
||||
TypeSafeDefaultModel = "jev-latest"
|
||||
)
|
||||
|
||||
// EffectiveAuditBackend returns openai or typesafe. Omitted or unknown values default to openai.
|
||||
func (h HitlConfig) EffectiveAuditBackend() string {
|
||||
switch strings.ToLower(strings.TrimSpace(h.AuditBackend)) {
|
||||
case HitlAuditBackendTypeSafe, "jev", "type-safe", "typesafe-ai":
|
||||
return HitlAuditBackendTypeSafe
|
||||
default:
|
||||
return HitlAuditBackendOpenAI
|
||||
}
|
||||
}
|
||||
|
||||
// TypeSafeConfigEffective returns TypeSafe endpoint settings. Empty base_url/model use defaults; API key is never inherited from the main OpenAI channel.
|
||||
func (h HitlConfig) TypeSafeConfigEffective() (baseURL, apiKey, model string) {
|
||||
baseURL = strings.TrimSpace(h.AuditModel.BaseURL)
|
||||
if baseURL == "" {
|
||||
baseURL = TypeSafeDefaultBaseURL
|
||||
}
|
||||
apiKey = strings.TrimSpace(h.AuditModel.APIKey)
|
||||
model = strings.TrimSpace(h.AuditModel.Model)
|
||||
if model == "" {
|
||||
model = TypeSafeDefaultModel
|
||||
}
|
||||
return strings.TrimSuffix(baseURL, "/"), apiKey, model
|
||||
}
|
||||
|
||||
// AuditModelEffective returns the audit-agent model config with empty fields inherited from the main model config.
|
||||
func (h HitlConfig) AuditModelEffective(main OpenAIConfig) OpenAIConfig {
|
||||
out := main
|
||||
@@ -1291,6 +1324,22 @@ func (c HitlConfig) EffectiveAuditAgentPromptForMode(mode string) string {
|
||||
return DefaultHitlAuditAgentPrompt()
|
||||
}
|
||||
|
||||
// JevOperatorPolicy returns a custom audit-strategy prompt for TypeSafe Jev.
|
||||
// Built-in default prompts stay encoded as Jev questions and are not copied into state.
|
||||
func (c HitlConfig) JevOperatorPolicy(mode string) string {
|
||||
effective := strings.TrimSpace(c.EffectiveAuditAgentPromptForMode(mode))
|
||||
var def string
|
||||
if normalizeHitlModeForPrompt(mode) == "review_edit" {
|
||||
def = strings.TrimSpace(DefaultHitlAuditAgentPromptReviewEdit())
|
||||
} else {
|
||||
def = strings.TrimSpace(DefaultHitlAuditAgentPrompt())
|
||||
}
|
||||
if effective == "" || effective == def {
|
||||
return ""
|
||||
}
|
||||
return effective
|
||||
}
|
||||
|
||||
func normalizeHitlModeForPrompt(mode string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(mode)) {
|
||||
case "review_edit":
|
||||
|
||||
@@ -75,6 +75,34 @@ func TestLoadIgnoresLegacyAuthPasswordField(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestHitlEffectiveAuditBackend(t *testing.T) {
|
||||
if got := (HitlConfig{}).EffectiveAuditBackend(); got != HitlAuditBackendOpenAI {
|
||||
t.Fatalf("empty backend = %q, want openai", got)
|
||||
}
|
||||
if got := (HitlConfig{AuditBackend: "Jev"}).EffectiveAuditBackend(); got != HitlAuditBackendTypeSafe {
|
||||
t.Fatalf("jev alias = %q, want typesafe", got)
|
||||
}
|
||||
if got := (HitlConfig{AuditBackend: "claude"}).EffectiveAuditBackend(); got != HitlAuditBackendOpenAI {
|
||||
t.Fatalf("unknown backend = %q, want openai", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHitlTypeSafeConfigEffectiveDoesNotInheritMainKey(t *testing.T) {
|
||||
gotURL, gotKey, gotModel := (HitlConfig{
|
||||
AuditBackend: "typesafe",
|
||||
AuditModel: OpenAIConfig{APIKey: "ts-key"},
|
||||
}).TypeSafeConfigEffective()
|
||||
if gotURL != TypeSafeDefaultBaseURL {
|
||||
t.Fatalf("base url = %q, want default", gotURL)
|
||||
}
|
||||
if gotKey != "ts-key" {
|
||||
t.Fatalf("api key = %q, want ts-key", gotKey)
|
||||
}
|
||||
if gotModel != TypeSafeDefaultModel {
|
||||
t.Fatalf("model = %q, want default", gotModel)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHitlAuditModelEffectiveFallsBackToMainConfig(t *testing.T) {
|
||||
main := OpenAIConfig{
|
||||
Provider: "openai",
|
||||
|
||||
@@ -29,3 +29,15 @@ func TestDefaultHitlAuditAgentPromptReviewEditKeepsEditedArguments(t *testing.T)
|
||||
t.Fatal("review-edit prompt must require a matched rule")
|
||||
}
|
||||
}
|
||||
|
||||
func TestJevOperatorPolicySkipsDefaultPrompt(t *testing.T) {
|
||||
if got := (HitlConfig{}).JevOperatorPolicy("approval"); got != "" {
|
||||
t.Fatalf("empty config should not send default prompt to Jev, got %q", got)
|
||||
}
|
||||
if got := (HitlConfig{AuditAgentPrompt: DefaultHitlAuditAgentPrompt()}).JevOperatorPolicy("approval"); got != "" {
|
||||
t.Fatalf("default prompt should not be sent to Jev, got %q", got)
|
||||
}
|
||||
if got := (HitlConfig{AuditAgentPrompt: "拦截所有命令执行"}).JevOperatorPolicy("approval"); got != "拦截所有命令执行" {
|
||||
t.Fatalf("custom prompt=%q", got)
|
||||
}
|
||||
}
|
||||
@@ -24,6 +24,7 @@ import (
|
||||
"cyberstrike-ai/internal/openai"
|
||||
"cyberstrike-ai/internal/security"
|
||||
"cyberstrike-ai/internal/toolguard"
|
||||
"cyberstrike-ai/internal/typesafe"
|
||||
|
||||
"github.com/cloudwego/eino/schema"
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -891,6 +892,7 @@ func (h *ConfigHandler) UpdateConfig(c *gin.Context) {
|
||||
}
|
||||
|
||||
if req.Hitl != nil {
|
||||
h.config.Hitl.AuditBackend = req.Hitl.EffectiveAuditBackend()
|
||||
h.config.Hitl.AuditModel = req.Hitl.AuditModel
|
||||
h.config.Hitl.ToolWhitelist = mergeHitlToolWhitelistSlice(nil, req.Hitl.ToolWhitelist)
|
||||
if strings.TrimSpace(req.Hitl.DefaultMode) != "" {
|
||||
@@ -911,6 +913,7 @@ func (h *ConfigHandler) UpdateConfig(c *gin.Context) {
|
||||
h.config.Hitl.RetentionDays = &v
|
||||
}
|
||||
h.logger.Info("更新HITL配置",
|
||||
zap.String("audit_backend", h.config.Hitl.AuditBackend),
|
||||
zap.String("default_reviewer", h.config.Hitl.DefaultReviewer),
|
||||
zap.Int("tool_whitelist", len(h.config.Hitl.ToolWhitelist)),
|
||||
)
|
||||
@@ -1317,6 +1320,61 @@ func (h *ConfigHandler) TestOpenAI(c *gin.Context) {
|
||||
})
|
||||
}
|
||||
|
||||
// TestTypeSafeRequest 测试 TypeSafe / Jev 连接。
|
||||
type TestTypeSafeRequest struct {
|
||||
BaseURL string `json:"base_url"`
|
||||
APIKey string `json:"api_key"`
|
||||
Model string `json:"model"`
|
||||
}
|
||||
|
||||
// TestTypeSafe 用一条最小 Noul 验证 TypeSafe System One 是否可用。
|
||||
func (h *ConfigHandler) TestTypeSafe(c *gin.Context) {
|
||||
var req TestTypeSafeRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "无效的请求参数: " + err.Error()})
|
||||
return
|
||||
}
|
||||
if strings.TrimSpace(req.APIKey) == "" {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "TypeSafe API Key 不能为空"})
|
||||
return
|
||||
}
|
||||
|
||||
client := typesafe.NewClient(req.BaseURL, req.APIKey, req.Model, nil)
|
||||
ctx, cancel := context.WithTimeout(c.Request.Context(), 30*time.Second)
|
||||
defer cancel()
|
||||
start := time.Now()
|
||||
result, err := client.SystemOne(ctx, "connectivity ping", map[string]typesafe.Question{
|
||||
"ok": typesafe.Noul("Is this a connectivity test ping?", "Yes, this is only a ping.", "No."),
|
||||
})
|
||||
if err != nil {
|
||||
if apiErr, ok := err.(*typesafe.APIError); ok {
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": false,
|
||||
"error": fmt.Sprintf("API 返回错误 (HTTP %d): %s", apiErr.StatusCode, apiErr.Body),
|
||||
"status_code": apiErr.StatusCode,
|
||||
})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": false,
|
||||
"error": "连接失败: " + err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
model := strings.TrimSpace(req.Model)
|
||||
if result != nil && strings.TrimSpace(result.Model) != "" {
|
||||
model = result.Model
|
||||
}
|
||||
if model == "" {
|
||||
model = config.TypeSafeDefaultModel
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"model": model,
|
||||
"latency_ms": time.Since(start).Milliseconds(),
|
||||
})
|
||||
}
|
||||
|
||||
// ListModelsRequest 获取模型列表请求(OpenAI 兼容 GET /models)。
|
||||
type ListModelsRequest struct {
|
||||
Provider string `json:"provider"`
|
||||
@@ -2149,6 +2207,7 @@ func (h *ConfigHandler) MergeHitlToolWhitelistIntoConfig(add []string) error {
|
||||
func updateHitlConfig(doc *yaml.Node, cfg config.HitlConfig) {
|
||||
root := doc.Content[0]
|
||||
hitlNode := ensureMap(root, "hitl")
|
||||
setStringInMap(hitlNode, "audit_backend", cfg.EffectiveAuditBackend())
|
||||
auditModelNode := ensureMap(hitlNode, "audit_model")
|
||||
setStringInMap(auditModelNode, "provider", cfg.AuditModel.Provider)
|
||||
setStringInMap(auditModelNode, "base_url", cfg.AuditModel.BaseURL)
|
||||
|
||||
@@ -659,10 +659,13 @@ func (h *AgentHandler) waitHITLApproval(runCtx context.Context, cancelRun contex
|
||||
expiresAt := approvalStartedAt.Add(cfg.Timeout)
|
||||
approvalExpiresAt = &expiresAt
|
||||
}
|
||||
auditBackend, auditModel := h.hitlAuditEngineInfo()
|
||||
payload["hitlApproval"] = map[string]interface{}{
|
||||
"createdAt": approvalStartedAt,
|
||||
"timeoutSeconds": timeoutSeconds,
|
||||
"expiresAt": approvalExpiresAt,
|
||||
"auditBackend": auditBackend,
|
||||
"auditModel": auditModel,
|
||||
}
|
||||
payloadRaw, _ := json.Marshal(payload)
|
||||
p, err := h.hitlManager.CreatePendingInterrupt(conversationID, assistantMessageID, cfg.Mode, toolName, toolCallID, string(payloadRaw), cfg.Reviewer)
|
||||
@@ -1072,11 +1075,14 @@ type setHitlDefaultConfigReq struct {
|
||||
}
|
||||
|
||||
func (h *AgentHandler) hitlDefaultConfigResponse() gin.H {
|
||||
backend, model := h.hitlAuditEngineInfo()
|
||||
return gin.H{
|
||||
"defaultMode": h.hitlEffectiveDefaultMode(),
|
||||
"defaultReviewer": h.hitlEffectiveDefaultReviewer(),
|
||||
"defaultTimeoutSeconds": h.hitlEffectiveDefaultTimeoutSeconds(),
|
||||
"hitlGlobalToolWhitelist": h.hitlConfigGlobalToolWhitelist(),
|
||||
"auditBackend": backend,
|
||||
"auditModel": model,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -10,7 +10,9 @@ import (
|
||||
"time"
|
||||
|
||||
"cyberstrike-ai/internal/config"
|
||||
"cyberstrike-ai/internal/hitl"
|
||||
"cyberstrike-ai/internal/openai"
|
||||
"cyberstrike-ai/internal/typesafe"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.uber.org/zap"
|
||||
@@ -23,6 +25,9 @@ func (h *AgentHandler) auditAgentReview(ctx context.Context, hitlMode, toolName
|
||||
return hitlDecision{Decision: "reject", Comment: "audit agent: handler unavailable"}
|
||||
}
|
||||
mode := normalizeHitlMode(hitlMode)
|
||||
if h.config != nil && h.config.Hitl.EffectiveAuditBackend() == config.HitlAuditBackendTypeSafe {
|
||||
return h.auditAgentReviewTypeSafe(ctx, mode, toolName, payload)
|
||||
}
|
||||
prompt := config.DefaultHitlAuditAgentPrompt()
|
||||
if h.config != nil {
|
||||
prompt = h.config.Hitl.EffectiveAuditAgentPromptForMode(mode)
|
||||
@@ -109,6 +114,34 @@ func (h *AgentHandler) auditLLMConfig() config.OpenAIConfig {
|
||||
return config.OpenAIConfig{}
|
||||
}
|
||||
|
||||
func (h *AgentHandler) auditAgentReviewTypeSafe(ctx context.Context, hitlMode, toolName string, payload map[string]interface{}) hitlDecision {
|
||||
if h == nil || h.config == nil {
|
||||
return hitlDecision{Decision: "reject", Comment: "audit agent: TypeSafe 未配置"}
|
||||
}
|
||||
baseURL, apiKey, model := h.config.Hitl.TypeSafeConfigEffective()
|
||||
if apiKey == "" {
|
||||
return hitlDecision{Decision: "reject", Comment: "audit agent: TypeSafe API Key 未配置"}
|
||||
}
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
callCtx, cancel := context.WithTimeout(ctx, 90*time.Second)
|
||||
defer cancel()
|
||||
|
||||
client := typesafe.NewClient(baseURL, apiKey, model, nil)
|
||||
policy := h.config.Hitl.JevOperatorPolicy(hitlMode)
|
||||
result, err := client.SystemOne(callCtx, hitl.BuildJevState(hitlMode, toolName, payload, policy), hitl.JevAuditQuestions(policy))
|
||||
if err != nil {
|
||||
h.logger.Warn("审计 Agent TypeSafe 调用失败", zap.Error(err), zap.String("tool", toolName))
|
||||
return hitlDecision{Decision: "reject", Comment: "audit agent: TypeSafe 调用失败,保守拒绝"}
|
||||
}
|
||||
decision, comment := hitl.DecideJev(result)
|
||||
if comment == "" {
|
||||
comment = "audit agent: " + decision
|
||||
}
|
||||
return hitlDecision{Decision: decision, Comment: comment}
|
||||
}
|
||||
|
||||
func buildAuditAgentReviewInput(hitlMode, toolName string, payload map[string]interface{}) string {
|
||||
review := map[string]interface{}{
|
||||
"hitlMode": normalizeHitlMode(hitlMode),
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"cyberstrike-ai/internal/config"
|
||||
)
|
||||
|
||||
func TestParseAuditAgentLLMContentApprove(t *testing.T) {
|
||||
@@ -65,6 +68,17 @@ func TestParseAuditAgentLLMContentWithEditedArguments(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuditAgentReviewTypeSafeMissingAPIKey(t *testing.T) {
|
||||
h := &AgentHandler{config: &config.Config{Hitl: config.HitlConfig{AuditBackend: "typesafe"}}}
|
||||
d := h.auditAgentReview(context.Background(), "approval", "exec", nil)
|
||||
if d.Decision != "reject" {
|
||||
t.Fatalf("decision=%s", d.Decision)
|
||||
}
|
||||
if !strings.Contains(d.Comment, "TypeSafe API Key") {
|
||||
t.Fatalf("comment=%s", d.Comment)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildAuditAgentReviewInputIncludesMode(t *testing.T) {
|
||||
s := buildAuditAgentReviewInput("review_edit", "execute", map[string]interface{}{
|
||||
"arguments": `{"command":"pwd"}`,
|
||||
|
||||
@@ -0,0 +1,74 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strings"
|
||||
|
||||
"cyberstrike-ai/internal/config"
|
||||
)
|
||||
|
||||
func (h *AgentHandler) hitlAuditEngineInfo() (backend, model string) {
|
||||
backend = config.HitlAuditBackendOpenAI
|
||||
if h == nil || h.config == nil {
|
||||
return backend, ""
|
||||
}
|
||||
backend = h.config.Hitl.EffectiveAuditBackend()
|
||||
if backend == config.HitlAuditBackendTypeSafe {
|
||||
_, _, model = h.config.Hitl.TypeSafeConfigEffective()
|
||||
return backend, model
|
||||
}
|
||||
return backend, strings.TrimSpace(h.config.Hitl.AuditModelEffective(h.config.OpenAI).Model)
|
||||
}
|
||||
|
||||
func stringifyHitlJSON(v any) string {
|
||||
if v == nil {
|
||||
return ""
|
||||
}
|
||||
if s, ok := v.(string); ok {
|
||||
return strings.TrimSpace(s)
|
||||
}
|
||||
b, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
var s string
|
||||
if json.Unmarshal(b, &s) == nil {
|
||||
return strings.TrimSpace(s)
|
||||
}
|
||||
return strings.TrimSpace(string(b))
|
||||
}
|
||||
|
||||
func inferHitlAuditBackendFromComment(comment string) string {
|
||||
c := strings.ToLower(comment)
|
||||
if strings.Contains(comment, "TypeSafe") || strings.Contains(comment, "破坏分") ||
|
||||
strings.Contains(c, "choice=") || strings.Contains(comment, "Jev") {
|
||||
return config.HitlAuditBackendTypeSafe
|
||||
}
|
||||
if strings.TrimSpace(comment) == "" {
|
||||
return ""
|
||||
}
|
||||
return config.HitlAuditBackendOpenAI
|
||||
}
|
||||
|
||||
func hitlAuditBackendFromRecord(decidedBy, comment, payloadJSON string) (backend, model string) {
|
||||
if normalizeHitlDecidedBy(decidedBy) != "audit_agent" {
|
||||
return "", ""
|
||||
}
|
||||
var root map[string]any
|
||||
if err := json.Unmarshal([]byte(payloadJSON), &root); err == nil {
|
||||
if appr, ok := root["hitlApproval"].(map[string]any); ok {
|
||||
raw := stringifyHitlJSON(appr["auditBackend"])
|
||||
if raw != "" {
|
||||
backend = (config.HitlConfig{AuditBackend: raw}).EffectiveAuditBackend()
|
||||
}
|
||||
model = stringifyHitlJSON(appr["auditModel"])
|
||||
}
|
||||
}
|
||||
if backend == "" {
|
||||
backend = inferHitlAuditBackendFromComment(comment)
|
||||
}
|
||||
if backend == "" {
|
||||
backend = config.HitlAuditBackendOpenAI
|
||||
}
|
||||
return backend, model
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"cyberstrike-ai/internal/config"
|
||||
)
|
||||
|
||||
func TestHitlAuditEngineInfoTypeSafe(t *testing.T) {
|
||||
h := &AgentHandler{config: &config.Config{
|
||||
OpenAI: config.OpenAIConfig{Model: "gpt-4o"},
|
||||
Hitl: config.HitlConfig{AuditBackend: "typesafe"},
|
||||
}}
|
||||
backend, model := h.hitlAuditEngineInfo()
|
||||
if backend != config.HitlAuditBackendTypeSafe {
|
||||
t.Fatalf("backend=%q", backend)
|
||||
}
|
||||
if model != config.TypeSafeDefaultModel {
|
||||
t.Fatalf("model=%q, want %s", model, config.TypeSafeDefaultModel)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHitlAuditEngineInfoOpenAIInheritsMainModel(t *testing.T) {
|
||||
h := &AgentHandler{config: &config.Config{
|
||||
OpenAI: config.OpenAIConfig{Model: "gpt-4o-mini"},
|
||||
Hitl: config.HitlConfig{AuditBackend: "openai"},
|
||||
}}
|
||||
backend, model := h.hitlAuditEngineInfo()
|
||||
if backend != config.HitlAuditBackendOpenAI {
|
||||
t.Fatalf("backend=%q", backend)
|
||||
}
|
||||
if model != "gpt-4o-mini" {
|
||||
t.Fatalf("model=%q", model)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHitlAuditBackendFromRecordPrefersPayload(t *testing.T) {
|
||||
backend, model := hitlAuditBackendFromRecord("audit_agent", "audit agent: 实际操作:探测", `{
|
||||
"hitlApproval": {"auditBackend": "typesafe", "auditModel": "jev-latest"}
|
||||
}`)
|
||||
if backend != config.HitlAuditBackendTypeSafe || model != "jev-latest" {
|
||||
t.Fatalf("backend=%q model=%q", backend, model)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHitlAuditBackendFromRecordInfersJevComment(t *testing.T) {
|
||||
backend, _ := hitlAuditBackendFromRecord("audit_agent",
|
||||
"audit agent: 未命中破坏性规则,默认放行;最高破坏分=破坏业务可用性 0.12;choice=approve(0.90)",
|
||||
`{}`)
|
||||
if backend != config.HitlAuditBackendTypeSafe {
|
||||
t.Fatalf("backend=%q", backend)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHitlAuditBackendFromRecordInfersOpenAIComment(t *testing.T) {
|
||||
backend, _ := hitlAuditBackendFromRecord("audit_agent",
|
||||
"audit agent: 实际操作:读取 /etc/passwd;命中规则:A3",
|
||||
`{}`)
|
||||
if backend != config.HitlAuditBackendOpenAI {
|
||||
t.Fatalf("backend=%q", backend)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHitlAuditBackendFromRecordIgnoresHuman(t *testing.T) {
|
||||
backend, model := hitlAuditBackendFromRecord("human", "人工通过", `{"hitlApproval":{"auditBackend":"typesafe"}}`)
|
||||
if backend != "" || model != "" {
|
||||
t.Fatalf("backend=%q model=%q", backend, model)
|
||||
}
|
||||
}
|
||||
@@ -56,6 +56,7 @@ func hitlInterruptRowToMap(
|
||||
if messageID.Valid {
|
||||
msgID = messageID.String
|
||||
}
|
||||
auditBackend, auditModel := hitlAuditBackendFromRecord(decidedBy, comment.String, payload)
|
||||
return map[string]interface{}{
|
||||
"id": id,
|
||||
"conversationId": cid,
|
||||
@@ -69,6 +70,8 @@ func hitlInterruptRowToMap(
|
||||
"decision": decision.String,
|
||||
"comment": comment.String,
|
||||
"decidedBy": decidedBy,
|
||||
"auditBackend": auditBackend,
|
||||
"auditModel": auditModel,
|
||||
"createdAt": createdAt,
|
||||
"decidedAt": func() interface{} {
|
||||
if decidedAt.Valid {
|
||||
|
||||
@@ -4913,6 +4913,50 @@ func (h *OpenAPIHandler) GetOpenAPISpec(c *gin.Context) {
|
||||
},
|
||||
},
|
||||
},
|
||||
"/api/config/test-typesafe": map[string]interface{}{
|
||||
"post": map[string]interface{}{
|
||||
"tags": []string{"配置管理"},
|
||||
"summary": "测试 TypeSafe Jev 连接",
|
||||
"description": "发送一条最小 Noul 请求,验证 TypeSafe System One API Key 是否可用。",
|
||||
"operationId": "testTypeSafe",
|
||||
"requestBody": map[string]interface{}{
|
||||
"required": true,
|
||||
"content": map[string]interface{}{
|
||||
"application/json": map[string]interface{}{
|
||||
"schema": map[string]interface{}{
|
||||
"type": "object",
|
||||
"required": []string{"api_key"},
|
||||
"properties": map[string]interface{}{
|
||||
"base_url": map[string]interface{}{"type": "string", "description": "可选,默认 https://api.typesafe.ai"},
|
||||
"api_key": map[string]interface{}{"type": "string", "description": "TypeSafe API Key"},
|
||||
"model": map[string]interface{}{"type": "string", "description": "可选,默认 jev-latest", "example": "jev-latest"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
"responses": map[string]interface{}{
|
||||
"200": map[string]interface{}{
|
||||
"description": "测试结果",
|
||||
"content": map[string]interface{}{
|
||||
"application/json": map[string]interface{}{
|
||||
"schema": map[string]interface{}{
|
||||
"type": "object",
|
||||
"properties": map[string]interface{}{
|
||||
"success": map[string]interface{}{"type": "boolean"},
|
||||
"error": map[string]interface{}{"type": "string"},
|
||||
"model": map[string]interface{}{"type": "string"},
|
||||
"latency_ms": map[string]interface{}{"type": "number"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
"400": map[string]interface{}{"description": "参数错误"},
|
||||
"401": map[string]interface{}{"description": "未授权"},
|
||||
},
|
||||
},
|
||||
},
|
||||
"/api/config/list-models": map[string]interface{}{
|
||||
"post": map[string]interface{}{
|
||||
"tags": []string{"配置管理"},
|
||||
|
||||
@@ -0,0 +1,195 @@
|
||||
package hitl
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"cyberstrike-ai/internal/typesafe"
|
||||
)
|
||||
|
||||
const (
|
||||
jevRejectThreshold = 0.55
|
||||
|
||||
jevQAvailability = "destroys_availability"
|
||||
jevQDataLoss = "irreversible_data_loss"
|
||||
jevQCredentials = "credential_or_acl_change"
|
||||
jevQConfig = "config_tamper"
|
||||
jevQPentest = "authorized_pentest_payload"
|
||||
jevQOperatorPolicy = "violates_operator_policy"
|
||||
jevQDecision = "decision"
|
||||
|
||||
jevDefaultPolicy = "Approve authorized pentest operations by default. Reject only when the call would clearly destroy live business availability, irreversibly destroy real data, change production credentials/ACLs, or tamper with live system configuration. If both a pentest payload and a destructive action are present, reject."
|
||||
)
|
||||
|
||||
const jevStateMaxFieldRunes = 8000
|
||||
|
||||
// JevAuditQuestions is the built-in CyberStrikeAI approval policy, asked in English for Jev accuracy.
|
||||
// A non-empty operatorPolicy adds a Noul/Choice overlay so the organization prompt is evaluated as state, not as a chat system prompt.
|
||||
func JevAuditQuestions(operatorPolicy string) map[string]typesafe.Question {
|
||||
qs := map[string]typesafe.Question{
|
||||
jevQAvailability: typesafe.Noul(
|
||||
"Would executing this tool call stop, disable, reboot, or overwhelm a live business service?",
|
||||
"The call would stop/kill/reboot a service, or generate DoS-level load on a production system.",
|
||||
"The call would not take the business offline. Recon, scans, and authorized exploit payloads count as no.",
|
||||
),
|
||||
jevQDataLoss: typesafe.Noul(
|
||||
"Would executing this tool call cause irreversible loss or mass destruction of real business data or files?",
|
||||
"DROP/TRUNCATE/DELETE of real tables, rm of important files, wiping a database, or similar irreversible destruction.",
|
||||
"No irreversible data destruction. Authorized pentest writes, webshell drops, and read/query operations count as no.",
|
||||
),
|
||||
jevQCredentials: typesafe.Noul(
|
||||
"Would executing this tool call change real user passwords, force users offline, or modify accounts/roles/permissions?",
|
||||
"Password reset, account/role/ACL changes, or forcing real users offline.",
|
||||
"No account-system tampering. Exploit payloads that do not change production identity systems count as no.",
|
||||
),
|
||||
jevQConfig: typesafe.Noul(
|
||||
"Would executing this tool call change server, service, firewall, startup, or scheduled-task configuration of a live system?",
|
||||
"Modifying system/Web/database/firewall/startup/cron configuration on a real host.",
|
||||
"No live configuration change. Writing a test file or exploit artifact is not configuration tampering.",
|
||||
),
|
||||
jevQPentest: typesafe.Noul(
|
||||
"Is this an authorized penetration-testing action such as recon, scanning, exploit payload, webshell upload, reverse shell, or credential spraying at a non-DoS rate?",
|
||||
"Typical authorized pentest activity that should be allowed even if it looks aggressive or performs writes.",
|
||||
"Not a pentest payload, or the action is clearly aimed at destroying business availability or data.",
|
||||
),
|
||||
jevQDecision: typesafe.Choice(
|
||||
"Should CyberStrikeAI approve this tool call? Approve authorized pentest activity. Reject only actions that would clearly destroy live business availability, irreversibly destroy real data, change production credentials/ACLs, or tamper with live system configuration.",
|
||||
map[string]string{
|
||||
"approve": "Safe to execute in an authorized pentest. Includes exploit payloads, webshells, reverse shells, scans, and recon.",
|
||||
"reject": "Would clearly damage live business availability, destroy real data, change production accounts, or alter live system configuration.",
|
||||
},
|
||||
),
|
||||
}
|
||||
if strings.TrimSpace(operatorPolicy) == "" {
|
||||
return qs
|
||||
}
|
||||
qs[jevQOperatorPolicy] = typesafe.Noul(
|
||||
"Would executing this tool call be forbidden by the operatorPolicy text in the state? Read operatorPolicy even if it is written in Chinese. Treat exec, execute, shell, bash, zsh, and any command argument as command execution.",
|
||||
"operatorPolicy forbids this action, for example by blocking all command execution or this specific tool or command.",
|
||||
"operatorPolicy does not forbid this action.",
|
||||
)
|
||||
qs[jevQDecision] = typesafe.Choice(
|
||||
"Should this tool call be approved? Follow operatorPolicy in the state first; it is the organization's written approval policy and takes precedence over default pentest-allow rules. If operatorPolicy forbids this action, reject. Still reject live-system destruction, irreversible data loss, credential/ACL changes, and live configuration tampering.",
|
||||
map[string]string{
|
||||
"approve": "operatorPolicy allows this action or is silent, and it would not destroy a live system.",
|
||||
"reject": "operatorPolicy forbids this action, or it would destroy live business availability, data, credentials, or configuration.",
|
||||
},
|
||||
)
|
||||
return qs
|
||||
}
|
||||
|
||||
// BuildJevState keeps only the fields Jev needs. Large cognition blobs are truncated to avoid context rot.
|
||||
func BuildJevState(hitlMode, toolName string, payload map[string]interface{}, operatorPolicy string) map[string]interface{} {
|
||||
policy := jevDefaultPolicy
|
||||
if strings.TrimSpace(operatorPolicy) != "" {
|
||||
policy = "Follow operatorPolicy first. It is the organization's written approval policy and may be in Chinese. If it forbids this action, reject. The built-in floor still rejects live-system destruction."
|
||||
}
|
||||
state := map[string]interface{}{
|
||||
"hitlMode": strings.TrimSpace(hitlMode),
|
||||
"toolName": strings.TrimSpace(toolName),
|
||||
"policy": policy,
|
||||
}
|
||||
if s := strings.TrimSpace(operatorPolicy); s != "" {
|
||||
state["operatorPolicy"] = truncateRunes(s, jevStateMaxFieldRunes)
|
||||
}
|
||||
if payload == nil {
|
||||
return state
|
||||
}
|
||||
for _, k := range []string{"arguments", "argumentsObj", "command", "userMessage"} {
|
||||
if v, ok := payload[k]; ok && v != nil && fmt.Sprint(v) != "" {
|
||||
state[k] = truncateJevValue(v)
|
||||
}
|
||||
}
|
||||
return state
|
||||
}
|
||||
|
||||
func truncateJevValue(v interface{}) interface{} {
|
||||
switch t := v.(type) {
|
||||
case string:
|
||||
return truncateRunes(t, jevStateMaxFieldRunes)
|
||||
case map[string]interface{}, []interface{}:
|
||||
b, err := json.Marshal(t)
|
||||
if err != nil {
|
||||
return truncateRunes(fmt.Sprint(t), jevStateMaxFieldRunes)
|
||||
}
|
||||
s := string(b)
|
||||
if len([]rune(s)) <= jevStateMaxFieldRunes {
|
||||
return t
|
||||
}
|
||||
return truncateRunes(s, jevStateMaxFieldRunes)
|
||||
default:
|
||||
return truncateRunes(fmt.Sprint(t), jevStateMaxFieldRunes)
|
||||
}
|
||||
}
|
||||
|
||||
func truncateRunes(s string, max int) string {
|
||||
rs := []rune(s)
|
||||
if max <= 0 || len(rs) <= max {
|
||||
return s
|
||||
}
|
||||
return string(rs[:max]) + "…"
|
||||
}
|
||||
|
||||
// DecideJev maps TypeSafe answers onto approve/reject using code-owned policy.
|
||||
// Destructive Nouls and a custom operator-policy Noul win over pentest Nouls.
|
||||
// Uncertain cases follow the product default: approve, unless a custom policy Choice rejects with high confidence.
|
||||
func DecideJev(result *typesafe.Result) (decision, comment string) {
|
||||
if result == nil {
|
||||
return "reject", "audit agent: TypeSafe 无有效响应,保守拒绝"
|
||||
}
|
||||
availability := result.Noul(jevQAvailability)
|
||||
dataLoss := result.Noul(jevQDataLoss)
|
||||
credentials := result.Noul(jevQCredentials)
|
||||
configTamper := result.Noul(jevQConfig)
|
||||
pentest := result.Noul(jevQPentest)
|
||||
policyNoul := result.Noul(jevQOperatorPolicy)
|
||||
choice, choiceConf := result.Choice(jevQDecision)
|
||||
_, hasOperatorPolicy := result.Answers[jevQOperatorPolicy]
|
||||
|
||||
type hit struct {
|
||||
id string
|
||||
label string
|
||||
noul float64
|
||||
}
|
||||
hits := []hit{
|
||||
{jevQAvailability, "破坏业务可用性", availability},
|
||||
{jevQDataLoss, "不可逆数据破坏", dataLoss},
|
||||
{jevQCredentials, "账号权限篡改", credentials},
|
||||
{jevQConfig, "改系统配置", configTamper},
|
||||
{jevQOperatorPolicy, "组织审批策略", policyNoul},
|
||||
}
|
||||
var fired []string
|
||||
maxReject := 0.0
|
||||
topLabel := ""
|
||||
for _, h := range hits {
|
||||
if h.noul > maxReject {
|
||||
maxReject = h.noul
|
||||
topLabel = h.label
|
||||
}
|
||||
if h.noul >= jevRejectThreshold {
|
||||
fired = append(fired, fmt.Sprintf("%s=%.2f", h.label, h.noul))
|
||||
}
|
||||
}
|
||||
|
||||
decision = "approve"
|
||||
reason := "未命中破坏性规则,默认放行"
|
||||
if maxReject >= jevRejectThreshold {
|
||||
decision = "reject"
|
||||
reason = "命中拒绝规则:" + strings.Join(fired, ";")
|
||||
} else if strings.EqualFold(choice, "reject") && choiceConf >= 0.85 && (hasOperatorPolicy || (maxReject < 0.35 && pentest < 0.5)) {
|
||||
decision = "reject"
|
||||
if hasOperatorPolicy {
|
||||
reason = fmt.Sprintf("Jev 按组织策略拒绝(choice=%.2f,策略分=%.2f)", choiceConf, policyNoul)
|
||||
} else {
|
||||
reason = fmt.Sprintf("Jev 高置信拒绝(choice=%.2f,最高破坏分=%.2f)", choiceConf, maxReject)
|
||||
}
|
||||
}
|
||||
|
||||
if decision == "approve" && topLabel != "" {
|
||||
reason = fmt.Sprintf("%s;最高破坏分=%s %.2f;渗透payload=%.2f", reason, topLabel, maxReject, pentest)
|
||||
}
|
||||
|
||||
comment = fmt.Sprintf("audit agent: %s;choice=%s(%.2f)", reason, choice, choiceConf)
|
||||
return decision, comment
|
||||
}
|
||||
@@ -0,0 +1,154 @@
|
||||
package hitl
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"cyberstrike-ai/internal/typesafe"
|
||||
)
|
||||
|
||||
func TestDecideJevRejectsDestructive(t *testing.T) {
|
||||
dec, comment := DecideJev(&typesafe.Result{Answers: map[string]map[string]any{
|
||||
jevQAvailability: {"noul": 0.92},
|
||||
jevQDataLoss: {"noul": 0.1},
|
||||
jevQCredentials: {"noul": 0.05},
|
||||
jevQConfig: {"noul": 0.04},
|
||||
jevQPentest: {"noul": 0.8},
|
||||
jevQDecision: {"choice": "approve", "confidence": 0.4},
|
||||
}})
|
||||
if dec != "reject" {
|
||||
t.Fatalf("decision=%s comment=%s", dec, comment)
|
||||
}
|
||||
if !strings.Contains(comment, "破坏业务可用性") {
|
||||
t.Fatalf("comment=%s", comment)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecideJevApprovesPentestPayload(t *testing.T) {
|
||||
dec, _ := DecideJev(&typesafe.Result{Answers: map[string]map[string]any{
|
||||
jevQAvailability: {"noul": 0.08},
|
||||
jevQDataLoss: {"noul": 0.06},
|
||||
jevQCredentials: {"noul": 0.04},
|
||||
jevQConfig: {"noul": 0.05},
|
||||
jevQPentest: {"noul": 0.97},
|
||||
jevQDecision: {"choice": "approve", "confidence": 0.9},
|
||||
}})
|
||||
if dec != "approve" {
|
||||
t.Fatalf("decision=%s", dec)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecideJevDestructiveWinsOverPentest(t *testing.T) {
|
||||
dec, _ := DecideJev(&typesafe.Result{Answers: map[string]map[string]any{
|
||||
jevQAvailability: {"noul": 0.12},
|
||||
jevQDataLoss: {"noul": 0.88},
|
||||
jevQCredentials: {"noul": 0.1},
|
||||
jevQConfig: {"noul": 0.1},
|
||||
jevQPentest: {"noul": 0.95},
|
||||
jevQDecision: {"choice": "approve", "confidence": 0.7},
|
||||
}})
|
||||
if dec != "reject" {
|
||||
t.Fatalf("decision=%s", dec)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecideJevUncertainApproves(t *testing.T) {
|
||||
dec, _ := DecideJev(&typesafe.Result{Answers: map[string]map[string]any{
|
||||
jevQAvailability: {"noul": 0.4},
|
||||
jevQDataLoss: {"noul": 0.2},
|
||||
jevQCredentials: {"noul": 0.1},
|
||||
jevQConfig: {"noul": 0.1},
|
||||
jevQPentest: {"noul": 0.3},
|
||||
jevQDecision: {"choice": "reject", "confidence": 0.5},
|
||||
}})
|
||||
if dec != "approve" {
|
||||
t.Fatalf("decision=%s", dec)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildJevStateOmitsCognitionBlobs(t *testing.T) {
|
||||
state := BuildJevState("approval", "exec", map[string]interface{}{
|
||||
"arguments": `{"command":"id"}`,
|
||||
"userMessage": "whoami",
|
||||
"thinking": "long chain",
|
||||
"reasoningChain": "should not appear",
|
||||
}, "")
|
||||
if state["toolName"] != "exec" {
|
||||
t.Fatalf("toolName=%v", state["toolName"])
|
||||
}
|
||||
if _, ok := state["thinking"]; ok {
|
||||
t.Fatal("thinking should be omitted")
|
||||
}
|
||||
if _, ok := state["reasoningChain"]; ok {
|
||||
t.Fatal("reasoningChain should be omitted")
|
||||
}
|
||||
if state["arguments"] != `{"command":"id"}` {
|
||||
t.Fatalf("arguments=%v", state["arguments"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestJevAuditQuestionsCoverPolicyAxes(t *testing.T) {
|
||||
qs := JevAuditQuestions("")
|
||||
for _, id := range []string{jevQAvailability, jevQDataLoss, jevQCredentials, jevQConfig, jevQPentest, jevQDecision} {
|
||||
if _, ok := qs[id]; !ok {
|
||||
t.Fatalf("missing question %s", id)
|
||||
}
|
||||
}
|
||||
if _, ok := qs[jevQOperatorPolicy]; ok {
|
||||
t.Fatal("default questions should not include operator policy overlay")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildJevStateIncludesOperatorPolicy(t *testing.T) {
|
||||
state := BuildJevState("approval", "exec", map[string]interface{}{"command": "id"}, "拦截所有命令执行")
|
||||
if state["operatorPolicy"] != "拦截所有命令执行" {
|
||||
t.Fatalf("operatorPolicy=%v", state["operatorPolicy"])
|
||||
}
|
||||
policy, _ := state["policy"].(string)
|
||||
if !strings.Contains(policy, "operatorPolicy") {
|
||||
t.Fatalf("policy=%v", state["policy"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestJevAuditQuestionsAddsPolicyOverlay(t *testing.T) {
|
||||
qs := JevAuditQuestions("拦截所有命令执行")
|
||||
if _, ok := qs[jevQOperatorPolicy]; !ok {
|
||||
t.Fatal("missing operator policy noul")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecideJevRejectsOperatorPolicy(t *testing.T) {
|
||||
dec, comment := DecideJev(&typesafe.Result{Answers: map[string]map[string]any{
|
||||
jevQAvailability: {"noul": 0.08},
|
||||
jevQDataLoss: {"noul": 0.06},
|
||||
jevQCredentials: {"noul": 0.04},
|
||||
jevQConfig: {"noul": 0.05},
|
||||
jevQPentest: {"noul": 0.01},
|
||||
jevQOperatorPolicy: {"noul": 0.91},
|
||||
jevQDecision: {"choice": "approve", "confidence": 0.2},
|
||||
}})
|
||||
if dec != "reject" {
|
||||
t.Fatalf("decision=%s comment=%s", dec, comment)
|
||||
}
|
||||
if !strings.Contains(comment, "组织审批策略") {
|
||||
t.Fatalf("comment=%s", comment)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecideJevPolicyChoiceRejectsEvenIfPentest(t *testing.T) {
|
||||
dec, comment := DecideJev(&typesafe.Result{Answers: map[string]map[string]any{
|
||||
jevQAvailability: {"noul": 0.1},
|
||||
jevQDataLoss: {"noul": 0.1},
|
||||
jevQCredentials: {"noul": 0.1},
|
||||
jevQConfig: {"noul": 0.1},
|
||||
jevQPentest: {"noul": 0.9},
|
||||
jevQOperatorPolicy: {"noul": 0.4},
|
||||
jevQDecision: {"choice": "reject", "confidence": 0.92},
|
||||
}})
|
||||
if dec != "reject" {
|
||||
t.Fatalf("decision=%s comment=%s", dec, comment)
|
||||
}
|
||||
if !strings.Contains(comment, "组织策略") {
|
||||
t.Fatalf("comment=%s", comment)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,196 @@
|
||||
package typesafe
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
DefaultBaseURL = "https://api.typesafe.ai"
|
||||
DefaultModel = "jev-latest"
|
||||
)
|
||||
|
||||
// Client calls TypeSafe System One (Jev).
|
||||
type Client struct {
|
||||
httpClient *http.Client
|
||||
baseURL string
|
||||
apiKey string
|
||||
model string
|
||||
}
|
||||
|
||||
// APIError is a non-2xx TypeSafe HTTP response.
|
||||
type APIError struct {
|
||||
StatusCode int
|
||||
Body string
|
||||
}
|
||||
|
||||
func (e *APIError) Error() string {
|
||||
return fmt.Sprintf("typesafe api error: status=%d body=%s", e.StatusCode, e.Body)
|
||||
}
|
||||
|
||||
// NewClient builds a System One client. Empty baseURL/model use TypeSafe defaults.
|
||||
func NewClient(baseURL, apiKey, model string, httpClient *http.Client) *Client {
|
||||
if httpClient == nil {
|
||||
httpClient = &http.Client{Timeout: 90 * time.Second}
|
||||
}
|
||||
baseURL = strings.TrimSuffix(strings.TrimSpace(baseURL), "/")
|
||||
if baseURL == "" {
|
||||
baseURL = DefaultBaseURL
|
||||
}
|
||||
model = strings.TrimSpace(model)
|
||||
if model == "" {
|
||||
model = DefaultModel
|
||||
}
|
||||
return &Client{
|
||||
httpClient: httpClient,
|
||||
baseURL: baseURL,
|
||||
apiKey: strings.TrimSpace(apiKey),
|
||||
model: model,
|
||||
}
|
||||
}
|
||||
|
||||
// Question is a typed System One question (noul / choice / score).
|
||||
type Question map[string]any
|
||||
|
||||
// Noul builds a yes/no question.
|
||||
func Noul(instructions string, trueMean, falseMean string) Question {
|
||||
q := Question{
|
||||
"type": "noul",
|
||||
"instructions": instructions,
|
||||
}
|
||||
if strings.TrimSpace(trueMean) != "" || strings.TrimSpace(falseMean) != "" {
|
||||
q["criteria"] = map[string]string{
|
||||
"true": trueMean,
|
||||
"false": falseMean,
|
||||
}
|
||||
}
|
||||
return q
|
||||
}
|
||||
|
||||
// Choice builds a closed-set question.
|
||||
func Choice(instructions string, criteria map[string]string) Question {
|
||||
return Question{
|
||||
"type": "choice",
|
||||
"instructions": instructions,
|
||||
"criteria": criteria,
|
||||
}
|
||||
}
|
||||
|
||||
// Result is a System One evaluation response.
|
||||
type Result struct {
|
||||
Model string `json:"model"`
|
||||
Answers map[string]map[string]any `json:"answers"`
|
||||
Usage Usage `json:"usage"`
|
||||
}
|
||||
|
||||
// Usage reports token counts.
|
||||
type Usage struct {
|
||||
InputTokens int `json:"input_tokens"`
|
||||
OutputTokens int `json:"output_tokens"`
|
||||
}
|
||||
|
||||
// Noul returns the probability that question id is yes.
|
||||
func (r *Result) Noul(id string) float64 {
|
||||
if r == nil {
|
||||
return 0
|
||||
}
|
||||
ans, ok := r.Answers[id]
|
||||
if !ok || ans == nil {
|
||||
return 0
|
||||
}
|
||||
switch v := ans["noul"].(type) {
|
||||
case float64:
|
||||
return v
|
||||
case json.Number:
|
||||
f, _ := v.Float64()
|
||||
return f
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
// Choice returns the selected option and confidence.
|
||||
func (r *Result) Choice(id string) (choice string, confidence float64) {
|
||||
if r == nil {
|
||||
return "", 0
|
||||
}
|
||||
ans, ok := r.Answers[id]
|
||||
if !ok || ans == nil {
|
||||
return "", 0
|
||||
}
|
||||
choice, _ = ans["choice"].(string)
|
||||
switch v := ans["confidence"].(type) {
|
||||
case float64:
|
||||
confidence = v
|
||||
case json.Number:
|
||||
confidence, _ = v.Float64()
|
||||
}
|
||||
return strings.TrimSpace(choice), confidence
|
||||
}
|
||||
|
||||
type systemOneRequest struct {
|
||||
State any `json:"state"`
|
||||
Model string `json:"model"`
|
||||
Questions map[string]Question `json:"questions"`
|
||||
}
|
||||
|
||||
// SystemOne evaluates state against questions.
|
||||
func (c *Client) SystemOne(ctx context.Context, state any, questions map[string]Question) (*Result, error) {
|
||||
if c == nil {
|
||||
return nil, fmt.Errorf("typesafe client is not initialized")
|
||||
}
|
||||
if strings.TrimSpace(c.apiKey) == "" {
|
||||
return nil, fmt.Errorf("typesafe api key is empty")
|
||||
}
|
||||
if len(questions) == 0 {
|
||||
return nil, fmt.Errorf("typesafe questions are empty")
|
||||
}
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
|
||||
body, err := json.Marshal(systemOneRequest{
|
||||
State: state,
|
||||
Model: c.model,
|
||||
Questions: questions,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("marshal typesafe payload: %w", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.baseURL+"/v1/systemone", bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("build typesafe request: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+c.apiKey)
|
||||
|
||||
resp, err := c.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("call typesafe api: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read typesafe response: %w", err)
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, &APIError{StatusCode: resp.StatusCode, Body: string(respBody)}
|
||||
}
|
||||
|
||||
var out Result
|
||||
if err := json.Unmarshal(respBody, &out); err != nil {
|
||||
return nil, fmt.Errorf("decode typesafe response: %w", err)
|
||||
}
|
||||
if out.Answers == nil {
|
||||
out.Answers = map[string]map[string]any{}
|
||||
}
|
||||
return &out, nil
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
package typesafe
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestSystemOneParsesNoulAndChoice(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/v1/systemone" {
|
||||
t.Fatalf("path = %s", r.URL.Path)
|
||||
}
|
||||
if got := r.Header.Get("Authorization"); got != "Bearer ts-key" {
|
||||
t.Fatalf("auth = %q", got)
|
||||
}
|
||||
var req systemOneRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if req.Model != "jev-latest" {
|
||||
t.Fatalf("model = %q", req.Model)
|
||||
}
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||
"model": "jev-1.13.0",
|
||||
"answers": map[string]any{
|
||||
"ok": map[string]any{"type": "noul", "noul": 0.91},
|
||||
"act": map[string]any{
|
||||
"type": "choice",
|
||||
"choice": "approve",
|
||||
"confidence": 0.8,
|
||||
},
|
||||
},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "ts-key", "", srv.Client())
|
||||
got, err := client.SystemOne(context.Background(), "ping", map[string]Question{
|
||||
"ok": Noul("Is this a ping?", "yes", "no"),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got.Noul("ok") != 0.91 {
|
||||
t.Fatalf("noul = %v", got.Noul("ok"))
|
||||
}
|
||||
choice, conf := got.Choice("act")
|
||||
if choice != "approve" || conf != 0.8 {
|
||||
t.Fatalf("choice=%q conf=%v", choice, conf)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemOneEmptyAPIKey(t *testing.T) {
|
||||
client := NewClient("", "", "", nil)
|
||||
if _, err := client.SystemOne(context.Background(), "ping", map[string]Question{"q": Noul("x", "", "")}); err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user