Skip to content
Open
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
2 changes: 1 addition & 1 deletion .github/workflows/ci.yml
Original file line number Diff line number Diff line change
Expand Up @@ -441,7 +441,7 @@ jobs:
run: |
# Untracked output is drift too: a new proto yields a new binding file,
# which `git diff` never reports.
test -z "$(git status --porcelain -- pkg/rpc core/types exts/guardrail/guardrail.pb.go exts/jev/jev.pb.go web/frontend/src/gen)"
test -z "$(git status --porcelain -- pkg/rpc core/decision core/types exts/guardrail/guardrail.pb.go exts/jev/jev.pb.go web/frontend/src/gen)"
# The AOP bindings are generated into the cyber-ui submodule, where the
# superproject tracks only the gitlink, so a regeneration inside it is
# invisible from here and the check has to run in the submodule.
Expand Down
4 changes: 2 additions & 2 deletions agent/provider/anthropic.go
Original file line number Diff line number Diff line change
Expand Up @@ -65,7 +65,7 @@ func (p *AnthropicProvider) ChatCompletion(ctx context.Context, req *ChatComplet
}
captureFrame(ctx, RawFrame{Provider: p.Name(), Protocol: ProviderAnthropic, Direction: "request", Transport: "http", Payload: bodyBytes, MediaType: "application/json"})

data, err := (&apiRequest{client: p.client, timeout: timeoutFromConfig(p.config.Timeout)}).do(
data, err := (&apiRequest{client: p.client, timeout: req.timeout(p.config.Timeout)}).do(
ctx, "POST", p.completionEndpoint(), bodyBytes, p.setAuthHeaders,
)
if err != nil {
Expand Down Expand Up @@ -100,7 +100,7 @@ func (p *AnthropicProvider) ChatCompletionStream(ctx context.Context, req *ChatC
}

parser := &anthropicStreamParser{}
events, err := streamSSE(ctx, p.client, timeoutFromConfig(p.config.Timeout),
events, err := streamSSE(ctx, p.client, req.timeout(p.config.Timeout),
p.completionEndpoint(), bodyBytes, p.setAuthHeaders, p.Name(), ProviderAnthropic,
false,
func(eventType string, data []byte) ([]ChatCompletionStreamEvent, error) {
Expand Down
43 changes: 21 additions & 22 deletions agent/provider/http.go
Original file line number Diff line number Diff line change
Expand Up @@ -26,11 +26,16 @@ func timeoutFromConfig(seconds int) time.Duration {
return time.Duration(seconds) * time.Second
}

func (r *ChatCompletionRequest) timeout(seconds int) time.Duration {
if r.Timeout > 0 {
return r.Timeout
}
return timeoutFromConfig(seconds)
}

func newHTTPClient(cfg *ProviderConfig) (*http.Client, error) {
timeout := timeoutFromConfig(cfg.Timeout)
transport := &http.Transport{
ResponseHeaderTimeout: timeout,
IdleConnTimeout: 90 * time.Second,
IdleConnTimeout: 90 * time.Second,
}
if cfg.Proxy != "" {
proxyURL, err := url.Parse(cfg.Proxy)
Expand Down Expand Up @@ -137,28 +142,33 @@ func streamSSE(
setHeaders(httpReq)
}

// The request owns its fallback, including the wait for response headers.
// A transport-wide deadline would cap longer background requests.
var stallDetected atomic.Bool
stallTimer := time.AfterFunc(timeout, func() {
stallDetected.Store(true)
reqCancel()
})
resp, err := client.Do(httpReq) //nolint:bodyclose // closed in goroutine below
if err != nil {
stallTimer.Stop()
reqCancel()
return nil, fmt.Errorf("http request: %w", err)
return nil, wrapReadError(ctx, stallDetected.Load(), timeout, "http request", err)
}

if resp.StatusCode < 200 || resp.StatusCode >= 300 {
defer stallTimer.Stop()
defer resp.Body.Close()
defer reqCancel()
respBody, timedOut, readErr := readAllWithCancelTimeout(resp.Body, reqCancel, timeout)
respBody, readErr := io.ReadAll(resp.Body)
if readErr != nil {
return nil, wrapReadError(ctx, timedOut, timeout, "read response", readErr)
return nil, wrapReadError(ctx, stallDetected.Load(), timeout, "read response", readErr)
}
captureFrame(ctx, RawFrame{Provider: providerName, Protocol: protocol, EventType: "error", Direction: "response", Transport: "sse", Payload: respBody, MediaType: "application/json"})
return nil, &APIError{StatusCode: resp.StatusCode, Message: string(respBody), Header: resp.Header.Clone()}
}

var stallDetected atomic.Bool
stallTimer := time.AfterFunc(timeout, func() {
stallDetected.Store(true)
reqCancel()
})
stallTimer.Reset(timeout)

events := make(chan ChatCompletionStreamEvent, 32)
go func() {
Expand Down Expand Up @@ -273,17 +283,6 @@ func wrapReadError(parentCtx context.Context, timedOut bool, timeout time.Durati
return fmt.Errorf("%s: %w", op, err)
}

func readAllWithCancelTimeout(r io.Reader, cancel context.CancelFunc, timeout time.Duration) ([]byte, bool, error) {
var timedOut atomic.Bool
timer := time.AfterFunc(timeout, func() {
timedOut.Store(true)
cancel()
})
defer timer.Stop()
body, err := io.ReadAll(r)
return body, timedOut.Load(), err
}

func clampInt(v, min, max, fallback int) int {
if v <= 0 {
return fallback
Expand Down
140 changes: 140 additions & 0 deletions agent/provider/jev/claim.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,140 @@
package jev

import (
"bytes"
"encoding/json"
"errors"
"fmt"
"io"
"math"
"slices"
"strings"

"github.com/chainreactors/cyber/core/decision"
)

type ClaimType = decision.ClaimType

const (
ClaimChoice = decision.ClaimType_choice
ClaimScore = decision.ClaimType_score
ClaimNoul = decision.ClaimType_noul
)

// Claim is a semantic judgment. Context supplies facts and option meanings;
// ordered options define the closed choice or score scale. Noul has no options.
// Persistence, evidence, compilation and execution are owned by consumers.
type Claim struct {
Type ClaimType `json:"type"`
Context string `json:"context"`
Options []string `json:"options,omitempty"`
}

func (c Claim) Validate() error {
if strings.TrimSpace(c.Context) == "" || len(c.Context) > 64<<10 {
return errors.New("invalid Claim context")
}
switch c.Type {
case ClaimChoice:
if len(c.Options) < 2 || len(c.Options) > 64 {
return errors.New("choice Claim needs two to sixty-four options")
}
case ClaimScore:
if len(c.Options) < 2 || len(c.Options) > 10 {
return errors.New("score Claim needs two to ten ordered levels")
}
case ClaimNoul:
if len(c.Options) != 0 {
return errors.New("noul Claim cannot have options")
}
default:
return errors.New("unsupported Claim type")
}
seen := map[string]bool{}
for _, option := range c.Options {
if strings.TrimSpace(option) == "" || len(option) > 1024 || seen[option] {
return errors.New("invalid or duplicate Claim option")
}
seen[option] = true
}
return nil
}

func (c Claim) Description() string {
if len(c.Options) == 0 {
return c.Type.String() + ": " + c.Context
}
return c.Type.String() + ": " + c.Context + "\nOptions (in order):\n" + strings.Join(c.Options, "\n")
}

func (c Claim) Proto() *decision.Claim {
return &decision.Claim{Type: c.Type, Context: c.Context, Options: slices.Clone(c.Options)}
}

func (c Claim) MarshalJSON() ([]byte, error) {
return json.Marshal(struct {
Type string `json:"type"`
Context string `json:"context"`
Options []string `json:"options,omitempty"`
}{c.Type.String(), c.Context, c.Options})
}

func (c *Claim) UnmarshalJSON(data []byte) error {
var value struct {
Type string `json:"type"`
Context string `json:"context"`
Options []string `json:"options,omitempty"`
}
d := json.NewDecoder(bytes.NewReader(data))
d.DisallowUnknownFields()
if err := d.Decode(&value); err != nil {
return err
}
if d.Decode(new(json.RawMessage)) != io.EOF {
return errors.New("extra Claim data")
}
kind, ok := decision.ClaimType_value[value.Type]
if !ok || kind == 0 {
return errors.New("unsupported Claim type")
}
claim := Claim{Type: ClaimType(kind), Context: value.Context, Options: value.Options}
if err := claim.Validate(); err != nil {
return err
}
*c = claim
return nil
}

// Evaluation is a closed result union; only one of choice, score or noul exists.
type Evaluation = decision.Evaluation

func (c Claim) Choice(e *Evaluation) (string, error) {
if e != nil && c.Type == ClaimChoice {
if v, ok := e.Value.(*decision.Evaluation_Choice); ok && slices.Contains(c.Options, v.Choice) {
return v.Choice, nil
}
}
return "", errors.New("invalid JEV choice binding")
}
func (c Claim) Score(e *Evaluation) (float64, error) {
if e != nil && c.Type == ClaimScore {
if v, ok := e.Value.(*decision.Evaluation_Score); ok {
return boundedNumber(v.Score, float64(len(c.Options)-1), "score")
}
}
return 0, errors.New("invalid JEV score response")
}
func (c Claim) Noul(e *Evaluation) (float64, error) {
if e != nil && c.Type == ClaimNoul {
if v, ok := e.Value.(*decision.Evaluation_Noul); ok {
return boundedNumber(v.Noul, 1, "noul")
}
}
return 0, errors.New("invalid JEV noul response")
}
func boundedNumber(value, upper float64, kind string) (float64, error) {
if math.IsNaN(value) || math.IsInf(value, 0) || value < 0 || value > upper {
return 0, fmt.Errorf("invalid JEV %s response", kind)
}
return value, nil
}
Loading
Loading