TokenEstimator.java 2.45 KB
package com.xly.service;

import dev.langchain4j.data.message.AiMessage;
import dev.langchain4j.data.message.ChatMessage;
import dev.langchain4j.data.message.SystemMessage;
import dev.langchain4j.data.message.ToolExecutionResultMessage;
import dev.langchain4j.data.message.UserMessage;
import dev.langchain4j.agent.tool.ToolExecutionRequest;

import java.util.List;

/**
 * 本地保守 token 估算器(不依赖任何模型请求)。
 *
 * <p>规则:CJK/全角字符 ≈ 1 token/字;其余字符 ≈ 3 字符/token。**故意高估**——
 * 高估只会少带几轮历史,低估会让 Ollama 从前面静默截头(system prompt 先死),两者代价不对称。
 * 校准回路:每次响应免费自带 prompt_eval_count,{@code TracingChatModelListener} 记录
 * 估算 vs 实际,实际逼近 num_ctx 即告警。
 */
public final class TokenEstimator {

    /** 每条消息的结构开销(role/分隔符等)。 */
    public static final int MSG_OVERHEAD = 8;

    private TokenEstimator() {
    }

    public static int estimate(String s) {
        if (s == null || s.isEmpty()) {
            return 0;
        }
        int cjk = 0;
        int other = 0;
        for (int i = 0; i < s.length(); i++) {
            if (s.charAt(i) >= 0x2E80) {
                cjk++;
            } else {
                other++;
            }
        }
        return cjk + (other + 2) / 3;
    }

    public static int estimate(ChatMessage m) {
        if (m == null) {
            return 0;
        }
        int n = MSG_OVERHEAD;
        if (m instanceof SystemMessage sm) {
            n += estimate(sm.text());
        } else if (m instanceof UserMessage um) {
            n += estimate(um.hasSingleText() ? um.singleText() : String.valueOf(um.contents()));
        } else if (m instanceof AiMessage am) {
            n += estimate(am.text());
            if (am.hasToolExecutionRequests()) {
                for (ToolExecutionRequest r : am.toolExecutionRequests()) {
                    n += estimate(r.name()) + estimate(r.arguments()) + 6;
                }
            }
        } else if (m instanceof ToolExecutionResultMessage tr) {
            n += estimate(tr.text()) + estimate(tr.toolName());
        } else {
            n += estimate(String.valueOf(m));
        }
        return n;
    }

    public static int estimate(List<ChatMessage> messages) {
        int n = 0;
        for (ChatMessage m : messages) {
            n += estimate(m);
        }
        return n;
    }
}