213 lines
4.9 KiB
Go
213 lines
4.9 KiB
Go
package ai
|
|
|
|
import (
|
|
"context"
|
|
"crypto/hmac"
|
|
"crypto/sha256"
|
|
"encoding/base64"
|
|
"encoding/json"
|
|
"fmt"
|
|
"net/http"
|
|
"net/url"
|
|
"resume-platform/pkg/config"
|
|
"strings"
|
|
"time"
|
|
|
|
"github.com/gorilla/websocket"
|
|
)
|
|
|
|
type SparkProvider struct {
|
|
cfg *config.SparkConfig
|
|
httpClient *http.Client
|
|
}
|
|
|
|
func NewSparkProvider(cfg *config.SparkConfig) *SparkProvider {
|
|
return &SparkProvider{
|
|
cfg: cfg,
|
|
httpClient: &http.Client{
|
|
Timeout: time.Duration(cfg.Timeout) * time.Second,
|
|
},
|
|
}
|
|
}
|
|
|
|
func (p *SparkProvider) GetName() string {
|
|
return "spark"
|
|
}
|
|
|
|
func (p *SparkProvider) Generate(ctx context.Context, prompt string, opts ...Option) (string, error) {
|
|
options := &Options{
|
|
MaxTokens: p.cfg.MaxTokens,
|
|
Temperature: 0.7,
|
|
TopP: 0.9,
|
|
}
|
|
for _, opt := range opts {
|
|
opt(options)
|
|
}
|
|
|
|
wsURL, err := p.buildAuthURL()
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
|
|
return p.sendWebSocketRequest(ctx, wsURL, prompt, options)
|
|
}
|
|
|
|
func (p *SparkProvider) buildAuthURL() (string, error) {
|
|
parsedURL, err := url.Parse(p.cfg.BaseURL)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
|
|
host := parsedURL.Host
|
|
path := parsedURL.Path
|
|
now := time.Now()
|
|
date := now.UTC().Format(time.RFC1123)
|
|
|
|
signatureOrigin := fmt.Sprintf("host: %s\ndate: %s\nGET %s HTTP/1.1", host, date, path)
|
|
|
|
mac := hmac.New(sha256.New, []byte(p.cfg.APISecret))
|
|
mac.Write([]byte(signatureOrigin))
|
|
signatureSha := base64.StdEncoding.EncodeToString(mac.Sum(nil))
|
|
|
|
authorizationOrigin := fmt.Sprintf(
|
|
`api_key="%s", algorithm="%s", headers="%s", signature="%s"`,
|
|
p.cfg.APIKey, "hmac-sha256", "host date request-line", signatureSha,
|
|
)
|
|
|
|
authorization := base64.StdEncoding.EncodeToString([]byte(authorizationOrigin))
|
|
|
|
query := url.Values{}
|
|
query.Set("authorization", authorization)
|
|
query.Set("date", date)
|
|
query.Set("host", host)
|
|
|
|
return p.cfg.BaseURL + "?" + query.Encode(), nil
|
|
}
|
|
|
|
func (p *SparkProvider) sendWebSocketRequest(ctx context.Context, wsURL string, prompt string, options *Options) (string, error) {
|
|
conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil)
|
|
if err != nil {
|
|
return "", fmt.Errorf("failed to dial websocket: %w", err)
|
|
}
|
|
defer conn.Close()
|
|
|
|
done := make(chan struct{})
|
|
var result strings.Builder
|
|
var lastErr error
|
|
|
|
go func() {
|
|
defer close(done)
|
|
for {
|
|
_, message, err := conn.ReadMessage()
|
|
if err != nil {
|
|
if websocket.IsCloseError(err, websocket.CloseNormalClosure) {
|
|
return
|
|
}
|
|
lastErr = fmt.Errorf("failed to read message: %w", err)
|
|
return
|
|
}
|
|
|
|
var resp map[string]interface{}
|
|
if err := json.Unmarshal(message, &resp); err != nil {
|
|
continue
|
|
}
|
|
|
|
if header, ok := resp["header"].(map[string]interface{}); ok {
|
|
if code, ok := header["code"].(float64); ok && code != 0 {
|
|
if msg, ok := header["message"].(string); ok {
|
|
lastErr = fmt.Errorf("API error: %s (code: %d)", msg, int(code))
|
|
} else {
|
|
lastErr = fmt.Errorf("API error with code: %d", int(code))
|
|
}
|
|
return
|
|
}
|
|
}
|
|
|
|
if payload, ok := resp["payload"].(map[string]interface{}); ok {
|
|
if choices, ok := payload["choices"].(map[string]interface{}); ok {
|
|
if text, ok := choices["text"].([]interface{}); ok && len(text) > 0 {
|
|
if item, ok := text[0].(map[string]interface{}); ok {
|
|
if content, ok := item["content"].(string); ok {
|
|
result.WriteString(content)
|
|
}
|
|
}
|
|
}
|
|
if status, ok := choices["status"].(float64); ok && status == 2 {
|
|
return
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}()
|
|
|
|
body := map[string]interface{}{
|
|
"header": map[string]interface{}{
|
|
"app_id": p.cfg.APPID,
|
|
"uid": "resume-platform",
|
|
},
|
|
"parameter": map[string]interface{}{
|
|
"chat": map[string]interface{}{
|
|
"domain": p.getDomainName(),
|
|
"temperature": options.Temperature,
|
|
"max_tokens": options.MaxTokens,
|
|
"top_k": 4,
|
|
"web_search": p.cfg.WebSearch,
|
|
"show_ref_label": p.cfg.ShowRefLabel,
|
|
},
|
|
},
|
|
"payload": map[string]interface{}{
|
|
"message": map[string]interface{}{
|
|
"text": []map[string]interface{}{
|
|
{
|
|
"role": "user",
|
|
"content": prompt,
|
|
},
|
|
},
|
|
},
|
|
},
|
|
}
|
|
|
|
jsonBody, err := json.Marshal(body)
|
|
if err != nil {
|
|
return "", fmt.Errorf("failed to marshal request: %w", err)
|
|
}
|
|
|
|
if err := conn.WriteMessage(websocket.TextMessage, jsonBody); err != nil {
|
|
return "", fmt.Errorf("failed to write message: %w", err)
|
|
}
|
|
|
|
select {
|
|
case <-done:
|
|
if lastErr != nil {
|
|
return "", lastErr
|
|
}
|
|
return result.String(), nil
|
|
case <-ctx.Done():
|
|
conn.Close()
|
|
return "", ctx.Err()
|
|
case <-time.After(time.Duration(p.cfg.Timeout) * time.Second):
|
|
conn.Close()
|
|
return "", fmt.Errorf("request timeout")
|
|
}
|
|
}
|
|
|
|
func (p *SparkProvider) getDomainName() string {
|
|
switch p.cfg.Model {
|
|
case "lite":
|
|
return "lite"
|
|
case "pro":
|
|
return "generalv3"
|
|
case "pro-128k":
|
|
return "pro-128k"
|
|
case "max":
|
|
return "generalv3.5"
|
|
case "max-32k":
|
|
return "max-32k"
|
|
case "ultra":
|
|
return "4.0Ultra"
|
|
case "kjwx":
|
|
return "kjwx"
|
|
default:
|
|
return "lite"
|
|
}
|
|
} |