231 lines
8.7 KiB
Go
231 lines
8.7 KiB
Go
package aimatching
|
||
|
||
import (
|
||
"bytes"
|
||
"context"
|
||
"encoding/json"
|
||
"errors"
|
||
"fmt"
|
||
"io"
|
||
"net/http"
|
||
"strings"
|
||
"time"
|
||
|
||
log "github.com/go-admin-team/go-admin-core/logger"
|
||
"github.com/google/uuid"
|
||
)
|
||
|
||
// Suggestion limits are enforced defensively here too, even though callers
|
||
// (e.g. shopeeproduct) already reject oversized requests with a domain
|
||
// specific message; this keeps the package safe for any future caller.
|
||
const (
|
||
maxSuggestSources = 100
|
||
maxSuggestCandidates = 150
|
||
)
|
||
|
||
// SuggestSource is one value that needs a suggestion, addressed by a
|
||
// request-scoped short id so the model can never emit anything but a member
|
||
// of the candidate set.
|
||
type SuggestSource struct {
|
||
ID string `json:"id"`
|
||
Label string `json:"label"`
|
||
}
|
||
|
||
// SuggestCandidate is one server-provided PDD spec value the model may
|
||
// choose, addressed the same way.
|
||
type SuggestCandidate struct {
|
||
ID string `json:"id"`
|
||
Label string `json:"label"`
|
||
}
|
||
|
||
// SuggestRequest asks for suggestions across many source values in a single
|
||
// call. Dimension only steers prompt wording ("color" or "size"); it is never
|
||
// sent as free text to identify anything beyond that.
|
||
type SuggestRequest struct {
|
||
Dimension string
|
||
ShopeeTitle string
|
||
PDDTitle string
|
||
Sources []SuggestSource
|
||
Candidates []SuggestCandidate
|
||
}
|
||
|
||
// SuggestDecision is the model's answer for one source. CandidateID is empty
|
||
// when the model had no reliable suggestion.
|
||
type SuggestDecision struct {
|
||
SourceID string
|
||
CandidateID string
|
||
Confidence float64
|
||
Reason string
|
||
}
|
||
|
||
// SuggestResult never contains a decision whose CandidateID is not one of the
|
||
// ids supplied in SuggestRequest.Candidates: SuggestBatch itself enforces that
|
||
// so callers cannot accidentally trust an out-of-band value.
|
||
type SuggestResult struct {
|
||
Decisions map[string]SuggestDecision
|
||
Provider string
|
||
Model string
|
||
}
|
||
|
||
// SuggestBatch generates draft suggestions only; it never writes to the
|
||
// database and never selects anything outside the supplied candidate set.
|
||
// Callers must still treat the result as an unsaved draft (#40, #46: AI 匹配
|
||
// 必须人工确认后生效).
|
||
func (s *Service) SuggestBatch(ctx context.Context, request SuggestRequest) (SuggestResult, error) {
|
||
if len(request.Sources) == 0 {
|
||
return SuggestResult{Decisions: map[string]SuggestDecision{}}, nil
|
||
}
|
||
if len(request.Sources) > maxSuggestSources || len(request.Candidates) > maxSuggestCandidates {
|
||
return SuggestResult{}, fail(CodeInvalidSetting, "AI 建议的候选或来源数量超过上限")
|
||
}
|
||
setting, apiKey, err := s.activeSetting(ctx)
|
||
if err != nil {
|
||
return SuggestResult{}, err
|
||
}
|
||
payload := openAIChatRequest{Model: setting.Model, Temperature: 0, Messages: []openAIMessage{
|
||
{Role: "system", Content: suggestSystemPrompt(request.Dimension)},
|
||
{Role: "user", Content: suggestPrompt(request)},
|
||
}}
|
||
body, err := json.Marshal(payload)
|
||
if err != nil {
|
||
return SuggestResult{}, &Error{Code: CodeProviderUnavailable, Message: "AI 建议请求生成失败", Cause: err}
|
||
}
|
||
ctx, cancel := context.WithTimeout(ctx, time.Duration(setting.TimeoutSeconds)*time.Second)
|
||
defer cancel()
|
||
httpRequest, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint(setting.BaseURL, "chat/completions"), bytes.NewReader(body))
|
||
if err != nil {
|
||
return SuggestResult{}, fail(CodeInvalidSetting, "AI 服务地址无效")
|
||
}
|
||
httpRequest.Header.Set("Authorization", "Bearer "+apiKey)
|
||
httpRequest.Header.Set("Content-Type", "application/json")
|
||
callID, startedAt := uuid.NewString(), time.Now()
|
||
response, err := s.httpClient().Do(httpRequest)
|
||
if err != nil {
|
||
s.logProviderFailure(ProviderFailureDiagnostic{CallID: callID, Operation: "suggest_batch", Kind: providerNetworkErrorKind(err), Duration: time.Since(startedAt)})
|
||
return SuggestResult{}, &Error{Code: CodeProviderUnavailable, Message: "AI 建议服务暂时不可用", Cause: err}
|
||
}
|
||
defer response.Body.Close()
|
||
limited := io.LimitReader(response.Body, 1<<20)
|
||
responseBody, readErr := io.ReadAll(limited)
|
||
if readErr != nil {
|
||
s.logProviderFailure(ProviderFailureDiagnostic{CallID: callID, Operation: "suggest_batch", Kind: "read_error", StatusCode: response.StatusCode, Duration: time.Since(startedAt)})
|
||
return SuggestResult{}, fail(CodeProviderUnavailable, "AI 建议服务暂时不可用")
|
||
}
|
||
if response.StatusCode < http.StatusOK || response.StatusCode >= http.StatusMultipleChoices {
|
||
s.logProviderFailure(ProviderFailureDiagnostic{CallID: callID, Operation: "suggest_batch", Kind: "http_status", StatusCode: response.StatusCode, Duration: time.Since(startedAt)})
|
||
return SuggestResult{}, fail(CodeProviderUnavailable, "AI 建议服务暂时不可用")
|
||
}
|
||
raw, err := parseSuggestChoices(responseBody)
|
||
if err != nil {
|
||
s.logProviderFailure(ProviderFailureDiagnostic{CallID: callID, Operation: "suggest_batch", Kind: "invalid_response", StatusCode: response.StatusCode, Duration: time.Since(startedAt)})
|
||
return SuggestResult{}, fail(CodeProviderUnavailable, "AI 建议响应无效")
|
||
}
|
||
|
||
sourceIDs := make(map[string]bool, len(request.Sources))
|
||
for _, source := range request.Sources {
|
||
sourceIDs[source.ID] = true
|
||
}
|
||
candidateIDs := make(map[string]bool, len(request.Candidates))
|
||
for _, candidate := range request.Candidates {
|
||
candidateIDs[candidate.ID] = true
|
||
}
|
||
|
||
decisions := make(map[string]SuggestDecision, len(raw))
|
||
for _, item := range raw {
|
||
sourceID := strings.TrimSpace(item.SourceID)
|
||
if sourceID == "" || !sourceIDs[sourceID] {
|
||
continue
|
||
}
|
||
if _, exists := decisions[sourceID]; exists {
|
||
// Duplicate source id in the model's answer: keep the first, the
|
||
// rest cannot be trusted to refer to the same intent.
|
||
continue
|
||
}
|
||
candidateID := strings.TrimSpace(item.CandidateID)
|
||
if candidateID != "" && !candidateIDs[candidateID] {
|
||
// Candidate id outside the supplied set: never guess, treat as no
|
||
// suggestion rather than translating it to anything.
|
||
candidateID = ""
|
||
}
|
||
confidence := 0.0
|
||
if item.Confidence != nil && *item.Confidence >= 0 && *item.Confidence <= 1 {
|
||
confidence = *item.Confidence
|
||
}
|
||
decisions[sourceID] = SuggestDecision{SourceID: sourceID, CandidateID: candidateID, Confidence: confidence, Reason: safeReason(item.Reason)}
|
||
}
|
||
return SuggestResult{Decisions: decisions, Provider: ProviderOpenAICompatible, Model: setting.Model}, nil
|
||
}
|
||
|
||
func (s *Service) logProviderFailure(diagnostic ProviderFailureDiagnostic) {
|
||
if s.ProviderFailureLogger != nil {
|
||
s.ProviderFailureLogger(diagnostic)
|
||
return
|
||
}
|
||
log.Warnf("AI provider call failed: call_id=%s operation=%s kind=%s status=%d duration_ms=%d",
|
||
diagnostic.CallID, diagnostic.Operation, diagnostic.Kind, diagnostic.StatusCode, diagnostic.Duration.Milliseconds())
|
||
}
|
||
|
||
func providerNetworkErrorKind(err error) string {
|
||
if errors.Is(err, context.DeadlineExceeded) {
|
||
return "timeout"
|
||
}
|
||
return "network_error"
|
||
}
|
||
|
||
func suggestSystemPrompt(dimension string) string {
|
||
noun := "颜色或尺码"
|
||
switch dimension {
|
||
case "color":
|
||
noun = "颜色"
|
||
case "size":
|
||
noun = "尺码"
|
||
}
|
||
return fmt.Sprintf(
|
||
"你负责为每个来源%s在给定候选中选择完全一致的原始标签,不得猜测、不得改写候选文字、不得选择候选之外的内容。"+
|
||
"每个来源最多对应一个候选,也可以没有可靠候选。只返回 JSON:"+
|
||
"{\"suggestions\":[{\"sourceId\":\"来源编号\",\"candidateId\":\"候选编号或空字符串\",\"confidence\":0到1之间的数字,\"reason\":\"简短原因\"}]},"+
|
||
"必须为每个来源都返回一条记录。", noun)
|
||
}
|
||
|
||
func suggestPrompt(request SuggestRequest) string {
|
||
payload := struct {
|
||
ShopeeTitle string `json:"shopeeTitle,omitempty"`
|
||
PDDTitle string `json:"pddTitle,omitempty"`
|
||
Sources []SuggestSource `json:"sources"`
|
||
Candidates []SuggestCandidate `json:"candidates"`
|
||
}{request.ShopeeTitle, request.PDDTitle, request.Sources, request.Candidates}
|
||
raw, _ := json.Marshal(payload)
|
||
return string(raw)
|
||
}
|
||
|
||
type suggestChoiceItem struct {
|
||
SourceID string `json:"sourceId"`
|
||
CandidateID string `json:"candidateId"`
|
||
Confidence *float64 `json:"confidence"`
|
||
Reason string `json:"reason"`
|
||
}
|
||
|
||
func parseSuggestChoices(raw []byte) ([]suggestChoiceItem, error) {
|
||
var response struct {
|
||
Choices []struct {
|
||
Message struct {
|
||
Content string `json:"content"`
|
||
} `json:"message"`
|
||
} `json:"choices"`
|
||
}
|
||
if err := json.Unmarshal(raw, &response); err != nil || len(response.Choices) == 0 {
|
||
return nil, errors.New("invalid provider response")
|
||
}
|
||
content := strings.TrimSpace(response.Choices[0].Message.Content)
|
||
content = strings.TrimPrefix(content, "```json")
|
||
content = strings.TrimPrefix(content, "```")
|
||
content = strings.TrimSuffix(strings.TrimSpace(content), "```")
|
||
var body struct {
|
||
Suggestions []suggestChoiceItem `json:"suggestions"`
|
||
}
|
||
if err := json.Unmarshal([]byte(content), &body); err != nil {
|
||
return nil, err
|
||
}
|
||
return body.Suggestions, nil
|
||
}
|