50 lines
1.6 KiB
Go
50 lines
1.6 KiB
Go
package chorus
|
|
|
|
import (
|
|
"context"
|
|
"encoding/base64"
|
|
"errors"
|
|
|
|
"git.ilapage.cn/OPC/chorus/internal/core/model"
|
|
"git.ilapage.cn/OPC/chorus/internal/core/provider"
|
|
)
|
|
|
|
const connectivityPrompt = "Chorus connectivity probe. Reply with OK."
|
|
|
|
type LiveProbe struct{ http provider.HTTPClient }
|
|
|
|
func NewLiveProbe(httpClient provider.HTTPClient) (*LiveProbe, error) {
|
|
if httpClient == nil {
|
|
return nil, errors.New("safe provider HTTP client is required")
|
|
}
|
|
return &LiveProbe{http: httpClient}, nil
|
|
}
|
|
|
|
func (p *LiveProbe) Probe(ctx context.Context, cfg ProbeConfiguration) error {
|
|
client, err := provider.NewOpenAI(p.http, provider.OpenAIConfig{
|
|
BaseURL: cfg.BaseURL, AuthType: cfg.AuthType, APIKey: cfg.APIKey,
|
|
ExtraBody: cfg.ExtraBody, MaxResponseBytes: cfg.MaxResponseBytes,
|
|
})
|
|
if err != nil {
|
|
return errors.New("prepare connectivity probe")
|
|
}
|
|
probeCtx, cancel := context.WithTimeout(ctx, cfg.Timeout)
|
|
defer cancel()
|
|
request := provider.Request{Kind: cfg.Kind, APIType: cfg.APIType, ModelID: cfg.ModelID, RenderedPrompt: connectivityPrompt}
|
|
if cfg.APIType == model.APIImagesEdits {
|
|
content, err := base64.StdEncoding.DecodeString("iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVQIHWP4z8DwHwAFgAI/ScL4WQAAAABJRU5ErkJggg==")
|
|
if err != nil {
|
|
return errors.New("prepare connectivity image")
|
|
}
|
|
defer clear(content)
|
|
request.Inputs = []provider.Input{{Content: content, MIMEType: "image/png", Role: model.RolePrimary, Position: 0}}
|
|
}
|
|
outputs, err := client.Generate(probeCtx, request)
|
|
if err != nil || len(outputs) == 0 {
|
|
return errors.New("connectivity probe failed")
|
|
}
|
|
return nil
|
|
}
|
|
|
|
var _ ProbeRunner = (*LiveProbe)(nil)
|