001/*
002 *  Copyright (c) 2023-2026, Agents-Flex (fuhai999@gmail.com).
003 *  <p>
004 *  Licensed under the Apache License, Version 2.0 (the "License");
005 *  you may not use this file except in compliance with the License.
006 *  You may obtain a copy of the License at
007 *  <p>
008 *  http://www.apache.org/licenses/LICENSE-2.0
009 *  <p>
010 *  Unless required by applicable law or agreed to in writing, software
011 *  distributed under the License is distributed on an "AS IS" BASIS,
012 *  WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
013 *  See the License for the specific language governing permissions and
014 *  limitations under the License.
015 */
016package com.agentsflex.core.util;
017
018import com.agentsflex.core.message.AiMessage;
019import com.agentsflex.core.message.Message;
020import com.agentsflex.core.message.ToolCall;
021import com.knuddels.jtokkit.Encodings;
022import com.knuddels.jtokkit.api.Encoding;
023import com.knuddels.jtokkit.api.EncodingType;
024
025import java.util.List;
026
027/**
028 * 静态工具类:更精确的本地 token 统计工具,模拟 OpenAI ChatCompletion 格式。
029 * 支持 function calls,按消息 role/name/内容序列化计数。
030 */
031public class LocalTokenCounter {
032
033    // 静态 Encoder,线程安全
034    private static Encoding ENCODING =
035        Encodings.newDefaultEncodingRegistry().getEncoding(EncodingType.CL100K_BASE);
036
037    public static void init(Encoding encoding) {
038        ENCODING = encoding;
039    }
040
041    /**
042     * 基于完整对话历史,为最后一条 AiMessage 计算并设置本地 token 字段。
043     *
044     * @param messages  完整的对话消息列表(按顺序,包含 system/user/ai)
045     * @param aiMessage 要设置 token 的 AiMessage(应为 messages 中的最后一条 assistant 消息)
046     */
047    public static void computeAndSetLocalTokens(List<Message> messages, AiMessage aiMessage) {
048        if (messages == null || messages.isEmpty() || aiMessage == null) {
049            return;
050        }
051
052        int promptTokens = countPromptTokens(messages);
053        int completionTokens = countCompletionTokens(aiMessage);
054
055        aiMessage.setLocalPromptTokens(promptTokens);
056        aiMessage.setLocalCompletionTokens(completionTokens);
057        aiMessage.setLocalTotalTokens(promptTokens + completionTokens);
058    }
059
060    /**
061     * 计算 prompt token(对话历史)
062     * 按 OpenAI ChatCompletion 格式,每条消息 role+content+name 固定 token
063     */
064    public static int countPromptTokens(List<? extends Message> messages) {
065        if (messages == null || messages.isEmpty()) return 0;
066
067        int total = 0;
068        for (Message msg : messages) {
069            total += countMessageTokens(msg);
070        }
071        // 结尾通常多一个 token,模拟 OpenAI 格式
072        total += 2;
073        return total;
074    }
075
076    /**
077     * 计算单条消息 token
078     */
079    private static int countMessageTokens(Message msg) {
080        int count = 0;
081
082        // role token
083        count += 1;
084
085        // content token
086        Object content = msg.getTextContent();
087        if (content != null) {
088            count += ENCODING.countTokens(content.toString());
089        }
090
091        return count;
092    }
093
094    /**
095     * 计算 AiMessage completion token
096     * 包含 fullContent / reasoningContent / functionCall
097     */
098    public static int countCompletionTokens(AiMessage aiMsg) {
099        if (aiMsg == null) return 0;
100
101        int count = 0;
102
103        // 生成的文本
104        if (aiMsg.getFullContent() != null) {
105            count += ENCODING.countTokens(aiMsg.getFullContent());
106        }
107
108        // 推理内容
109        if (aiMsg.getFullReasoningContent() != null) {
110            count += ENCODING.countTokens(aiMsg.getFullReasoningContent());
111        } else if (aiMsg.getReasoningContent() != null) {
112            count += ENCODING.countTokens(aiMsg.getReasoningContent());
113        }
114
115        // function call(按 JSON 序列化计算)
116        List<ToolCall> toolCalls = aiMsg.getToolCalls();
117        if (toolCalls != null && !toolCalls.isEmpty()) {
118            for (ToolCall toolCall : toolCalls) {
119                String serialized = toolCall.toJsonString();
120                count += ENCODING.countTokens(serialized);
121            }
122        }
123
124        // completion 固定尾部 token
125        count += 2;
126        return count;
127    }
128
129}