diff --git a/sdk/go/classifier.go b/sdk/go/classifier.go index b3363f4..d9babc2 100644 --- a/sdk/go/classifier.go +++ b/sdk/go/classifier.go @@ -178,3 +178,106 @@ func Classify(ctx context.Context, inputs, labels []string) ([]Result, error) { } return res.Results, nil } + +// Dimension configures one classification dimension. Labels is required; +// Instructions is optional guidance for ambiguous labels. +type Dimension struct { + Labels []string `json:"labels"` + Instructions string `json:"instructions,omitempty"` +} + +// DimensionRequest is the body of POST /v1/classify with dimensions. +type DimensionRequest struct { + Items []string `json:"items"` + Dimensions map[string]Dimension `json:"dimensions"` + Tier string `json:"tier,omitempty"` +} + +// DimensionFieldResult is one field in a multi-dimensional classification. +type DimensionFieldResult struct { + Label string `json:"label"` + Confidence *float64 `json:"confidence"` + Scores map[string]float64 `json:"scores"` + Model string `json:"model"` + MS int `json:"ms"` + Escalated bool `json:"escalated"` + Unscored string `json:"unscored,omitempty"` +} + +// DimensionResult is one item's classification across all requested dimensions. +type DimensionResult struct { + Dimensions map[string]DimensionFieldResult `json:"dimensions"` +} + +// DimensionResponse is the body of a successful dimensions call. +type DimensionResponse struct { + Tier string `json:"tier"` + Model string `json:"model"` + ModelsUsed []string `json:"modelsUsed"` + Results []DimensionResult `json:"results"` + Usage struct { + Items int `json:"items"` + Dimensions int `json:"dimensions"` + Classifications int `json:"classifications"` + Fallback int `json:"fallback"` + Escalated int `json:"escalated"` + MS int `json:"ms"` + } `json:"usage"` +} + +// ClassifyDimensions classifies items across multiple independent dimensions +// in one call. Up to 20 dimensions and 1,000 item × dimension decisions. +func (c *Client) ClassifyDimensions(ctx context.Context, req DimensionRequest) (*DimensionResponse, error) { + if len(req.Items) == 0 || len(req.Items) > 1000 { + return nil, errors.New("classifier.dev: items must hold 1 to 1,000 texts") + } + if len(req.Dimensions) == 0 || len(req.Dimensions) > 20 { + return nil, errors.New("classifier.dev: dimensions must hold 1 to 20 entries") + } + base := c.BaseURL + if base == "" { + base = "https://classifier.dev" + } + hc := c.HTTPClient + if hc == nil { + hc = http.DefaultClient + } + body, err := json.Marshal(req) + if err != nil { + return nil, err + } + httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, base+"/v1/classify", bytes.NewReader(body)) + if err != nil { + return nil, err + } + httpReq.Header.Set("Content-Type", "application/json") + httpReq.Header.Set("User-Agent", "classifier-dev-go/0.1.0") + if c.APIKey != "" { + httpReq.Header.Set("Authorization", "Bearer "+c.APIKey) + } + res, err := hc.Do(httpReq) + if err != nil { + return nil, err + } + defer res.Body.Close() + raw, err := io.ReadAll(res.Body) + if err != nil { + return nil, err + } + if res.StatusCode != http.StatusOK { + e := &Error{Status: res.StatusCode, Message: fmt.Sprintf("HTTP %d", res.StatusCode), Code: "http_" + strconv.Itoa(res.StatusCode)} + _ = json.Unmarshal(raw, e) + if s, err := strconv.Atoi(res.Header.Get("Retry-After")); err == nil { + e.RetryAfter = time.Duration(s) * time.Second + } + return nil, e + } + var out DimensionResponse + if err := json.Unmarshal(raw, &out); err != nil { + return nil, fmt.Errorf("classifier.dev: bad response: %w", err) + } + if len(out.Results) != len(req.Items) { + return nil, fmt.Errorf("classifier.dev: %d results for %d items", len(out.Results), len(req.Items)) + } + return &out, nil +} diff --git a/sdk/go/classifier_test.go b/sdk/go/classifier_test.go index 91704c9..b7c493d 100644 --- a/sdk/go/classifier_test.go +++ b/sdk/go/classifier_test.go @@ -76,3 +76,93 @@ func TestSingleResultShape(t *testing.T) { t.Fatalf("expected a shape error, got %v", err) } } + +func TestClassifyDimensions(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var body map[string]interface{} + _ = json.NewDecoder(r.Body).Decode(&body) + if _, ok := body["labels"]; ok { + t.Error("dimensions request must not carry labels") + } + items := body["items"].([]interface{}) + dims := body["dimensions"].(map[string]interface{}) + results := make([]map[string]interface{}, len(items)) + for i := range items { + fields := map[string]interface{}{} + for name, dim := range dims { + d := dim.(map[string]interface{}) + labels := d["labels"].([]interface{}) + conf := 0.85 + fields[name] = map[string]interface{}{ + "label": labels[0], "confidence": conf, + "scores": map[string]interface{}{labels[0].(string): conf}, + "model": "jev-test", "ms": 50, + } + } + results[i] = map[string]interface{}{"dimensions": fields} + } + out := map[string]interface{}{ + "tier": "fast", "model": "jev-test", "modelsUsed": []string{"jev-test"}, + "results": results, "usage": map[string]interface{}{ + "items": len(items), "dimensions": len(dims), + "classifications": len(items) * len(dims), + }, + } + _ = json.NewEncoder(w).Encode(out) + })) + defer srv.Close() + + c := &Client{BaseURL: srv.URL} + res, err := c.ClassifyDimensions(context.Background(), DimensionRequest{ + Items: []string{"checkout broke"}, + Dimensions: map[string]Dimension{ + "team": {Labels: []string{"billing", "platform"}}, + "kind": {Labels: []string{"bug", "request"}, Instructions: "bug = broken"}, + }, + }) + if err != nil { + t.Fatal(err) + } + if len(res.Results) != 1 { + t.Fatalf("expected 1 result, got %d", len(res.Results)) + } + if res.Results[0].Dimensions["team"].Label != "billing" { + t.Errorf("expected billing, got %s", res.Results[0].Dimensions["team"].Label) + } + if res.Results[0].Dimensions["kind"].Label != "bug" { + t.Errorf("expected bug, got %s", res.Results[0].Dimensions["kind"].Label) + } +} + +func TestClassifyDimensionsValidation(t *testing.T) { + cases := []struct { + name string + req DimensionRequest + msg string + }{ + {"no items", DimensionRequest{Dimensions: map[string]Dimension{"a": {Labels: []string{"x", "y"}}}}, "items must hold 1 to 1,000 texts"}, + {"no dimensions", DimensionRequest{Items: []string{"x"}}, "dimensions must hold 1 to 20 entries"}, + {"empty dimensions", DimensionRequest{Items: []string{"x"}, Dimensions: map[string]Dimension{}}, "dimensions must hold 1 to 20 entries"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + _, err := (&Client{}).ClassifyDimensions(context.Background(), tc.req) + if err == nil || err.Error() != "classifier.dev: "+tc.msg { + t.Fatalf("expected %q, got %v", tc.msg, err) + } + }) + } +} + +func TestClassifyDimensionsResultCountMismatch(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write([]byte(`{"results":[]}`)) + })) + defer srv.Close() + _, err := (&Client{BaseURL: srv.URL}).ClassifyDimensions(context.Background(), DimensionRequest{ + Items: []string{"x"}, Dimensions: map[string]Dimension{"a": {Labels: []string{"x", "y"}}}, + }) + if err == nil { + t.Fatal("expected error for result count mismatch") + } +}