package oci import ( "encoding/json" "fmt" "strings" "github.com/oracle/oci-go-sdk/v65/generativeaiinference" "oci-portal/internal/aiwire" ) // isCohereModel 依据模型名前缀判定 COHERE 系(走专属请求格式)。 func isCohereModel(name string) bool { return strings.HasPrefix(strings.ToLower(name), "cohere.") } // irToCohereSDK 把 IR 转为 COHERE 聊天请求;工具与结构化输出做有损降级 // (参数仅取 JSON Schema 顶层 properties,tool_choice / json_schema 的 name、strict 无对应)。 func irToCohereSDK(ir aiwire.ChatRequest, stream bool) (generativeaiinference.CohereChatRequest, error) { var req generativeaiinference.CohereChatRequest if err := cohereRejectUnsupported(ir); err != nil { return req, err } p, err := cohereSplitMessages(ir.Messages) if err != nil { return req, err } req = generativeaiinference.CohereChatRequest{ Message: &p.message, ChatHistory: p.history, ToolResults: p.toolResults, Tools: cohereTools(ir.Tools), ResponseFormat: cohereResponseFormat(ir.ResponseFormat), MaxTokens: ir.MaxTokens, Temperature: ir.Temperature, TopP: ir.TopP, TopK: ir.TopK, FrequencyPenalty: ir.FrequencyPenalty, PresencePenalty: ir.PresencePenalty, Seed: ir.Seed, } if p.preamble != "" { req.PreambleOverride = &p.preamble } if len(ir.Stop) > 0 { req.StopSequences = ir.Stop } if stream { req.IsStream = &stream includeUsage := ir.StreamOptions == nil || ir.StreamOptions.IncludeUsage req.StreamOptions = &generativeaiinference.StreamOptions{IsIncludeUsage: &includeUsage} } return req, nil } // cohereRejectUnsupported 拒绝 COHERE 无法承接的能力(多模态图片)。 func cohereRejectUnsupported(ir aiwire.ChatRequest) error { for _, m := range ir.Messages { for _, p := range m.Content.Parts { if p.Type == "image_url" { return fmt.Errorf("cohere 系模型暂不支持图片输入,请改用 meta/google 等多模态模型") } } } return nil } // cohereParts 是 IR 消息拆装为 COHERE 请求的中间结果。 type cohereParts struct { message string history []generativeaiinference.CohereMessage preamble string toolResults []generativeaiinference.CohereToolResult } // cohereSplitMessages 拆装消息:system 拼 preamble,末尾 user 抽为 message, // assistant(含 toolCalls)与其余 user 进 history,tool 结果转顶层 toolResults // (工具结果在末尾时 message 留空,COHERE 以 toolResults 续跑)。 func cohereSplitMessages(msgs []aiwire.ChatMessage) (cohereParts, error) { lastUser, lastTool := -1, -1 for i, m := range msgs { switch m.Role { case "user": lastUser = i case "tool": lastTool = i } } if lastUser == -1 && lastTool == -1 { return cohereParts{}, fmt.Errorf("cohere 系模型至少需要一条 user 消息") } var p cohereParts var preamble []string calls := map[string]generativeaiinference.CohereToolCall{} for i, m := range msgs { text := m.Content.JoinText() switch m.Role { case "system", "developer": preamble = append(preamble, text) case "assistant": p.history = append(p.history, cohereBotMessage(text, m.ToolCalls, calls)) case "tool": p.toolResults = append(p.toolResults, cohereToolResult(m, calls)) default: // user if i == lastUser && lastUser > lastTool { p.message = text continue } p.history = append(p.history, generativeaiinference.CohereUserMessage{Message: &text}) } } p.preamble = strings.Join(preamble, "\n") return p, nil } // cohereBotMessage 转 assistant 消息;toolCalls 同步登记进 id→call 映射供 tool 结果回查。 func cohereBotMessage(text string, tcs []aiwire.ToolCall, calls map[string]generativeaiinference.CohereToolCall) generativeaiinference.CohereChatBotMessage { msg := generativeaiinference.CohereChatBotMessage{} if text != "" { msg.Message = &text } for _, tc := range tcs { name := tc.Function.Name var params interface{} if json.Unmarshal([]byte(tc.Function.Arguments), ¶ms) != nil || params == nil { params = map[string]any{} } call := generativeaiinference.CohereToolCall{Name: &name, Parameters: ¶ms} calls[tc.ID] = call msg.ToolCalls = append(msg.ToolCalls, call) } return msg } // cohereToolResult 把 IR tool 消息转顶层工具结果;COHERE 工具调用无 id, // 靠 assistant 历史登记的映射回查,缺失时以 tool_call_id 名义调用兜底。 func cohereToolResult(m aiwire.ChatMessage, calls map[string]generativeaiinference.CohereToolCall) generativeaiinference.CohereToolResult { call, ok := calls[m.ToolCallID] if !ok { name := m.ToolCallID var params interface{} = map[string]any{} call = generativeaiinference.CohereToolCall{Name: &name, Parameters: ¶ms} } text := m.Content.JoinText() var output interface{} if json.Unmarshal([]byte(text), &output) != nil || output == nil { output = map[string]any{"output": text} } if arr, isArr := output.([]interface{}); isArr { return generativeaiinference.CohereToolResult{Call: &call, Outputs: arr} } if _, isMap := output.(map[string]interface{}); !isMap { output = map[string]any{"output": output} } return generativeaiinference.CohereToolResult{Call: &call, Outputs: []interface{}{output}} } // cohereTools 把 JSON Schema 工具定义降级为 COHERE 扁平参数表(嵌套结构有损:仅取顶层)。 func cohereTools(tools []aiwire.Tool) []generativeaiinference.CohereTool { if len(tools) == 0 { return nil } out := make([]generativeaiinference.CohereTool, 0, len(tools)) for _, t := range tools { name, desc := t.Function.Name, t.Function.Description if desc == "" { desc = name // Description 为 COHERE 必填 } out = append(out, generativeaiinference.CohereTool{ Name: &name, Description: &desc, ParameterDefinitions: cohereParams(t.Function.Parameters), }) } return out } // cohereParams 取 JSON Schema 顶层 properties 转扁平参数定义;嵌套 schema 只保留类型名。 func cohereParams(schema json.RawMessage) map[string]generativeaiinference.CohereParameterDefinition { var s struct { Properties map[string]struct { Type string `json:"type"` Description string `json:"description"` } `json:"properties"` Required []string `json:"required"` } if json.Unmarshal(schema, &s) != nil || len(s.Properties) == 0 { return nil } req := map[string]bool{} for _, r := range s.Required { req[r] = true } out := make(map[string]generativeaiinference.CohereParameterDefinition, len(s.Properties)) for k, v := range s.Properties { typ, desc := v.Type, v.Description if typ == "" { typ = "string" } pd := generativeaiinference.CohereParameterDefinition{Type: &typ} if desc != "" { pd.Description = &desc } if req[k] { t := true pd.IsRequired = &t } out[k] = pd } return out } // cohereResponseFormat 映射 json_object / json_schema(name、strict 无对应,有损降级)。 func cohereResponseFormat(rf *aiwire.ResponseFormat) generativeaiinference.CohereResponseFormat { if rf == nil { return nil } switch rf.Type { case "json_object", "json_schema": out := generativeaiinference.CohereResponseJsonFormat{} var in struct { Schema json.RawMessage `json:"schema"` } if json.Unmarshal(rf.JSONSchema, &in) == nil && len(in.Schema) > 0 { var schema interface{} if json.Unmarshal(in.Schema, &schema) == nil { out.Schema = &schema } } return out } return nil } // cohereSDKToIR 把 COHERE 非流式响应转回 IR;工具调用补派生 ID,存在时结束原因归 tool_calls。 func cohereSDKToIR(resp generativeaiinference.CohereChatResponse, model string) *aiwire.ChatResponse { msg := aiwire.ChatMessage{Role: "assistant", Content: aiwire.NewTextContent(deref(resp.Text))} for i, tc := range resp.ToolCalls { msg.ToolCalls = append(msg.ToolCalls, cohereCallToIR(deref(tc.Name), tc.Parameters, i)) } finish := mapFinishReason(string(resp.FinishReason)) if len(msg.ToolCalls) > 0 { finish = "tool_calls" } return &aiwire.ChatResponse{ Object: "chat.completion", Model: model, Choices: []aiwire.Choice{{Message: msg, FinishReason: finish}}, Usage: sdkUsageToIR(resp.Usage), } } // cohereCallToIR 生成 IR 工具调用;COHERE 无调用 id,按名称+序号派生(回传时按此回查)。 func cohereCallToIR(name string, params *interface{}, idx int) aiwire.ToolCall { args := "{}" if params != nil && *params != nil { if b, err := json.Marshal(*params); err == nil { args = string(b) } } return aiwire.ToolCall{ ID: fmt.Sprintf("call_%s_%d", name, idx), Type: "function", Function: aiwire.FunctionCall{Name: name, Arguments: args}, } }