package main
import (
"bytes"
"encoding/json"
"flag"
"fmt"
"io"
"log"
"net/http"
"os"
"sync"
"time"
amqp "github.com/rabbitmq/amqp091-go"
)
// ===================== 全局配置 =====================
const (
MQQueue = "client_ark_queue"
MQUrl = "amqp://guest:guest@127.0.0.1:5672/"
ArkModel = "doubao-seed-2-1-pro-260628"
ArkBaseUrl = "https://ark.cn-beijing.volces.com/api/v3/chat/completions"
)
var (
// Get API Key: https://ark.volcengine.com/region:cn-beijing/apikey
ArkApiKey = os.Getenv("ARK_API_KEY")
MockMode = false
)
// FlowConfig 流控配置
type FlowConfig struct {
CurrQPS int
MaxQPS int
Step int
Interval time.Duration
Count int
mu sync.Mutex
}
// 全局流控对象
var flow = &FlowConfig{
CurrQPS: 5,
MaxQPS: 30,
Step: 2,
Interval: 5 * time.Second,
Count: 0,
}
// ===================== 1. 极简斜率可控流控【核心】=====================
// startFlowSchedule 后台协程:匀速增长消费速率,控制斜率
func startFlowSchedule() {
ticker := time.NewTicker(flow.Interval)
defer ticker.Stop()
for range ticker.C {
flow.mu.Lock()
if flow.CurrQPS < flow.MaxQPS {
flow.CurrQPS += flow.Step
// 确保不超过最大值
if flow.CurrQPS > flow.MaxQPS {
flow.CurrQPS = flow.MaxQPS
}
log.Printf("QPS Limit Increased: Current=%d, Max=%d", flow.CurrQPS, flow.MaxQPS)
}
flow.mu.Unlock()
}
}
// startResetCount 后台协程:每秒重置请求计数器,精准控QPS
func startResetCount() {
ticker := time.NewTicker(1 * time.Second)
defer ticker.Stop()
for range ticker.C {
flow.mu.Lock()
flow.Count = 0
flow.mu.Unlock()
}
}
// allow 流控校验:是否允许发起本次请求
func allow() bool {
flow.mu.Lock()
defer flow.mu.Unlock()
if flow.Count < flow.CurrQPS {
flow.Count++
return true
}
return false
}
// ===================== 3. 方舟API调用【大模型请求核心】=====================
type ArkRequest struct {
Model string `json:"model"`
Messages []Message `json:"messages"`
}
type Message struct {
Role string `json:"role"`
Content string `json:"content"`
}
type ArkResponse struct {
Choices []struct {
Message Message `json:"message"`
} `json:"choices"`
Error *struct {
Message string `json:"message"`
Type string `json:"type"`
} `json:"error,omitempty"`
}
// callArk 调用方舟大模型API
func callArk(msg string) (string, error) {
if MockMode {
time.Sleep(50 * time.Millisecond) // 模拟网络延迟
return "Mocked Response: " + msg, nil
}
reqBody := ArkRequest{
Model: ArkModel,
Messages: []Message{
{Role: "user", Content: msg},
},
}
jsonData, err := json.Marshal(reqBody)
if err != nil {
return "", err
}
req, err := http.NewRequest("POST", ArkBaseUrl, bytes.NewBuffer(jsonData))
if err != nil {
return "", err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+ArkApiKey)
client := &http.Client{Timeout: 60 * time.Second}
resp, err := client.Do(req)
if err != nil {
return "", err
}
defer resp.Body.Close()
body, err := io.ReadAll(resp.Body)
if err != nil {
return "", err
}
if resp.StatusCode != http.StatusOK {
return "", fmt.Errorf("API error: status=%d, body=%s", resp.StatusCode, string(body))
}
var arkResp ArkResponse
if err := json.Unmarshal(body, &arkResp); err != nil {
return "", err
}
if arkResp.Error != nil {
return "", fmt.Errorf("API error: %s", arkResp.Error.Message)
}
if len(arkResp.Choices) > 0 {
return arkResp.Choices[0].Message.Content, nil
}
return "", fmt.Errorf("empty choices in response")
}
// ===================== 4. MQ消费者【填谷+斜率控速+调用方舟,核心主逻辑】=====================
// startConsumer 消费端核心:削峰后的流量,斜率可控匀速消费,调用方舟大模型
func startConsumer() {
conn, err := amqp.Dial(MQUrl)
if err != nil {
log.Fatalf("Failed to connect to RabbitMQ: %v", err)
}
defer conn.Close()
ch, err := conn.Channel()
if err != nil {
log.Fatalf("Failed to open a channel: %v", err)
}
defer ch.Close()
// 声明队列 (Durable=true)
q, err := ch.QueueDeclare(
MQQueue, // name
true, // durable
false, // delete when unused
false, // exclusive
false, // no-wait
nil, // arguments
)
if err != nil {
log.Fatalf("Failed to declare a queue: %v", err)
}
// 设置QoS (PrefetchCount=1)
err = ch.Qos(
1, // prefetch count
0, // prefetch size
false, // global
)
if err != nil {
log.Fatalf("Failed to set QoS: %v", err)
}
msgs, err := ch.Consume(
q.Name, // queue
"", // consumer
false, // auto-ack (设置为false,手动ACK)
false, // exclusive
false, // no-local
false, // no-wait
nil, // args
)
if err != nil {
log.Fatalf("Failed to register a consumer: %v", err)
}
log.Printf("消费启动|初始QPS:%d 最大QPS:%d 斜率可控", flow.CurrQPS, flow.MaxQPS)
forever := make(chan bool)
go func() {
for d := range msgs {
// 解析消息
var reqData map[string]string
if err := json.Unmarshal(d.Body, &reqData); err != nil {
log.Printf("Error parsing JSON: %v", err)
d.Nack(false, false) // 无法解析,丢弃或放入死信队列
continue
}
msgContent := reqData["msg"]
// 流控校验
if !allow() {
// 不通过则消息重回队列
d.Nack(false, true)
// 稍微sleep一下避免空转太快
time.Sleep(10 * time.Millisecond)
continue
}
// 调用方舟大模型
respContent, err := callArk(msgContent)
if err != nil {
log.Printf("调用方舟失败: %v", err)
d.Nack(false, true) // 失败重试
} else {
// 打印日志 (截取部分长度)
reqSnippet := msgContent
if len(reqSnippet) > 15 {
reqSnippet = reqSnippet[:15]
}
respSnippet := respContent
if len(respSnippet) > 20 {
respSnippet = respSnippet[:20]
}
log.Printf("请求:%s | 响应:%s", reqSnippet, respSnippet)
// 消费成功手动ACK
d.Ack(false)
}
}
}()
<-forever
}
// ===================== 2. MQ生产者【Client端请求投递,削峰】=====================
// startProducer 投递Client请求到MQ,无速率限制,承接峰值流量
func startProducer(count int) {
conn, err := amqp.Dial(MQUrl)
if err != nil {
log.Fatalf("Failed to connect to RabbitMQ: %v", err)
}
defer conn.Close()
ch, err := conn.Channel()
if err != nil {
log.Fatalf("Failed to open a channel: %v", err)
}
defer ch.Close()
q, err := ch.QueueDeclare(
MQQueue, // name
true, // durable
false, // delete when unused
false, // exclusive
false, // no-wait
nil, // arguments
)
if err != nil {
log.Fatalf("Failed to declare a queue: %v", err)
}
log.Printf("开始投递 %d 条消息...", count)
for i := 0; i < count; i++ {
msgContent := fmt.Sprintf("大模型请求%d: 技术方案总结", i)
reqData := map[string]string{"msg": msgContent}
body, _ := json.Marshal(reqData)
err = ch.Publish(
"", // exchange
q.Name, // routing key
false, // mandatory
false, // immediate
amqp.Publishing{
DeliveryMode: amqp.Persistent,
ContentType: "application/json",
Body: body,
})
if err != nil {
log.Fatalf("Failed to publish a message: %v", err)
}
}
log.Printf("成功投递 %d 条消息", count)
}
// ===================== 启动入口 =====================
func main() {
mode := flag.String("mode", "consumer", "运行模式: consumer | producer")
count := flag.Int("count", 100, "生产者发送消息数量")
mock := flag.Bool("mock", false, "是否开启Mock模式(不真实调用API)")
flag.Parse()
if *mock {
MockMode = true
log.Println("Mock模式已开启")
}
if *mode == "producer" {
startProducer(*count)
return
}
// 启动流控后台协程
go startFlowSchedule()
go startResetCount()
// 启动消费端
startConsumer()
}