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.model.client;
017
018import com.agentsflex.core.message.AiMessage;
019import com.agentsflex.core.model.chat.BaseChatModel;
020import com.agentsflex.core.model.chat.ChatContext;
021import com.agentsflex.core.model.chat.ChatContextHolder;
022import com.agentsflex.core.model.chat.StreamResponseListener;
023import com.agentsflex.core.model.chat.response.AiMessageResponse;
024import com.agentsflex.core.model.client.impl.SseClient;
025import com.agentsflex.core.parser.AiMessageParser;
026import com.agentsflex.core.parser.impl.DefaultAiMessageParser;
027import com.agentsflex.core.util.LocalTokenCounter;
028import com.agentsflex.core.util.Retryer;
029import com.agentsflex.core.util.StringUtil;
030import com.alibaba.fastjson2.JSON;
031import com.alibaba.fastjson2.JSONException;
032import com.alibaba.fastjson2.JSONObject;
033
034/**
035 * OpenAI 专用聊天客户端。
036 * <p>
037 * 封装了 HTTP 同步调用和 SSE 流式调用的具体实现。
038 */
039public class OpenAIChatClient extends ChatClient {
040
041    protected HttpClient httpClient;
042    protected AiMessageParser<JSONObject> aiMessageParser;
043
044    public OpenAIChatClient(BaseChatModel<?> chatModel) {
045        super(chatModel);
046    }
047
048    public HttpClient getHttpClient() {
049        if (httpClient == null) {
050            httpClient = new HttpClient();
051        }
052        return httpClient;
053    }
054
055    public void setHttpClient(HttpClient httpClient) {
056        this.httpClient = httpClient;
057    }
058
059    public StreamClient getStreamClient() {
060        // SseClient 默认实现是每次请求需要新建一个 SseClient 对象,方便进行 stop 调用
061        return new SseClient();
062    }
063
064
065    public AiMessageParser<JSONObject> getAiMessageParser() {
066        if (aiMessageParser == null) {
067            aiMessageParser = DefaultAiMessageParser.getOpenAIMessageParser();
068        }
069        return aiMessageParser;
070    }
071
072    public void setAiMessageParser(AiMessageParser<JSONObject> aiMessageParser) {
073        this.aiMessageParser = aiMessageParser;
074    }
075
076
077    @Override
078    public AiMessageResponse chat() {
079        HttpClient httpClient = getHttpClient();
080        ChatContext context = ChatContextHolder.currentContext();
081        ChatRequestSpec requestSpec = context.getRequestSpec();
082
083        String response = requestSpec.getRetryCount() > 0 ? Retryer.retry(() -> httpClient.post(requestSpec.getUrl(),
084            requestSpec.getHeaders(),
085            requestSpec.getBody()), requestSpec.getRetryCount(), requestSpec.getRetryInitialDelayMs())
086            : httpClient.post(requestSpec.getUrl(), requestSpec.getHeaders(), requestSpec.getBody());
087
088        if (StringUtil.noText(response)) {
089            return AiMessageResponse.error(context, response, "no content for response.");
090        }
091        try {
092            return parseResponse(response, context);
093        } catch (JSONException e) {
094            return AiMessageResponse.error(context, response, "invalid json response.");
095        }
096    }
097
098
099    protected AiMessageResponse parseResponse(String response, ChatContext context) {
100        JSONObject jsonObject = JSON.parseObject(response);
101        JSONObject error = jsonObject.getJSONObject("error");
102
103        AiMessageResponse messageResponse;
104        if (error != null && !error.isEmpty()) {
105            String message = error.getString("message");
106            messageResponse = AiMessageResponse.error(context, response, message);
107            messageResponse.setErrorType(error.getString("type"));
108            messageResponse.setErrorCode(error.getString("code"));
109        } else {
110            AiMessage aiMessage = getAiMessageParser().parse(jsonObject, context);
111            LocalTokenCounter.computeAndSetLocalTokens(context.getPrompt().getMessages(), aiMessage);
112            messageResponse = new AiMessageResponse(context, response, aiMessage);
113        }
114        return messageResponse;
115    }
116
117
118    @Override
119    public void chatStream(StreamResponseListener listener) {
120        StreamClient streamClient = getStreamClient();
121        ChatContext context = ChatContextHolder.currentContext();
122        StreamClientListener clientListener = new BaseStreamClientListener(
123            chatModel,
124            context,
125            streamClient,
126            listener,
127            getAiMessageParser()
128        );
129
130        ChatRequestSpec requestSpec = context.getRequestSpec();
131        if (requestSpec.getRetryCount() > 0) {
132            Retryer.retry(() -> streamClient.start(requestSpec.getUrl(), requestSpec.getHeaders(), requestSpec.getBody()
133                    , clientListener, chatModel.getConfig())
134                , requestSpec.getRetryCount()
135                , requestSpec.getRetryInitialDelayMs());
136        } else {
137            streamClient.start(requestSpec.getUrl(), requestSpec.getHeaders(), requestSpec.getBody()
138                , clientListener, chatModel.getConfig());
139        }
140    }
141
142
143}