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
103 changes: 103 additions & 0 deletions sdk/go/classifier.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
90 changes: 90 additions & 0 deletions sdk/go/classifier_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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")
}
}