#include "stdafx.h" #include "SH_AI.h" #include "AskAIPaletteSet.h" // 初始化 libcurl static SHAI::initCurl g_initCurl; namespace AIJSON { using namespace rapidjson; std::string w2utf(LPCTSTR s, DWORD cp) { return (LPCSTR)CW2A(s, cp); } std::string w2utf(const std::wstring &s, DWORD cp) { return w2utf(s.c_str(), cp); } std::wstring utf2w(LPCSTR s, DWORD cp) { return (LPCTSTR)CA2W(s, cp); } std::wstring utf2w(const std::string &s, DWORD cp) { return utf2w(s.c_str(), cp); } Document8::AllocatorType &getJsonAllocator() { static Document8 g_doc; return g_doc.GetAllocator(); } Value8 &addMember(Value8 &v, LPCSTR name, Value8 &value) { if (!name) { static Value8 vTmp; return vTmp; } if (!v.IsObject()) v.SetObject(); MIt8 it = v.FindMember(name); if (it != v.MemberEnd()) { it->value.Swap(value); } else { Value8 vName; vName.SetString(name, getJsonAllocator()); v.AddMember(vName, value, getJsonAllocator()); } return v; } Value8 &addMember(Value8 &v, LPCSTR name, LPCSTR value) { Value8 vv; vv.SetString(value, getJsonAllocator()); return addMember(v, name, vv); } Value8 &addMember(Value8 &v, LPCSTR name, const int &value) { Value8 vv(value); return addMember(v, name, vv); } Value8 &addMember(Value8 &v, LPCSTR name, const bool &value) { Value8 vv(value); return addMember(v, name, vv); } Value8 &addMember(Value8 &v, LPCSTR name, const float &value) { Value8 vv(value); return addMember(v, name, vv); } Value8 &pushBack(Value8 &arr, Value8 &value) { if (!arr.IsArray()) arr.SetArray(); arr.PushBack(value, getJsonAllocator()); return arr; } Value8 &pushBack(Value8 &arr, LPCSTR value) { Value8 v(StringRef(value)); return pushBack(arr, v); } Value8 &pushBack(Value8 &arr, const int &value) { Value8 vv(value); return pushBack(arr, vv); } Value8 &pushBack(Value8 &arr, const bool &value) { Value8 vv(value); return pushBack(arr, vv); } Value8 &pushBack(Value8 &arr, const float &value) { Value8 vv(value); return pushBack(arr, vv); } std::string write2Json(const Value8 &v, bool pretty) { StringBuffer8 buffer; if (pretty) { PrettyWriter8 writer(buffer); v.Accept(writer); } else { Writer8 writer(buffer); v.Accept(writer); } return buffer.GetString(); } bool json2File(LPCTSTR path, const Value8 &v, bool pretty) { CFile file(path, CFile::modeCreate | CFile::modeWrite | CFile::typeBinary); if (file.m_hFile) { std::string s(write2Json(v, pretty)); file.Write(s.c_str(), (UINT)(sizeof(char) * s.length())); file.Close(); } return false; } long getFileSize(FILE *fp) { if (!fp) return 0L; fseek(fp, 0L, SEEK_END); long fSize = ftell(fp); fseek(fp, 0L, 0L); return fSize; } bool file2Json(LPCTSTR pszJson, Document8 &doc) { if (!pszJson) return false; if (_taccess(pszJson, 0x0) != 0x0) return false; FILE *fp(NULL); char *readBuffer(NULL); CFile file(pszJson, CFile::typeBinary | CFile::modeRead); if (file.m_hFile) { ULONGLONG fSize = file.GetLength() + 1; char *pBuffer = new char[fSize]; ::memset(pBuffer, 0, sizeof(char) * fSize); file.Read(pBuffer, (UINT)fSize); file.Close(); doc.Parse(pBuffer); delete[] pBuffer; return !doc.HasParseError(); } return false; } } using namespace AIJSON; namespace { size_t split_string( const std::string &input , const std::string &delimiter , std::vector &result ) { if (input.empty()) return result.size(); // 输入为空直接返回 if (delimiter.empty()) { // 分隔符为空时返回原字符串 result.push_back(input); return result.size(); } size_t start = 0; size_t end = input.find(delimiter); while (end != std::string::npos) { // 截取从 start 到 end 的子字符串 result.push_back(input.substr(start, end - start)); // 跳过分隔符,更新起始位置 start = end + delimiter.length(); // 查找下一个分隔符位置 end = input.find(delimiter, start); } // 添加最后一个分段(剩余部分) result.push_back(input.substr(start)); return result.size(); } size_t request_WriteCallback(const char* pData, size_t length, void* pUserData) { SHAI::Response *response = (SHAI::Response*)pUserData; // 原封不动地使用您之前的字符串处理逻辑 std::string ss(pData, length); if (ss.find_first_of("data:") != 0) { response->set_finish_reason(ss.c_str()); return length; } std::string sResponse = ss; std::vector vs; split_string(sResponse, "data:", vs); for (size_t i = 0; i < vs.size(); i++) { if (!vs[i].empty()) response->readResponse(vs[i]); } return length; } size_t get_model_list_WriteCallback(const char *pData, size_t length, void *pUserData) { std::string *response = (std::string *)pUserData; std::string ss(pData, length); response->append(ss); return length; } } #pragma region Request std::string SHAI::Request::requestBody() const { Value8 body; addMember(body, "model", _model.c_str()); // 构建 message Value8 msg; msg.SetArray(); for (size_t i(0); i < _messages.size(); ++i) { Value8 m; addMember(m, "role", getRole(_messages[i].role).c_str()); addMember(m, "content", _messages[i].content.c_str()); pushBack(msg, m); } addMember(body, "messages", msg); addMember(body, "max_tokens", _max_tokens); addMember(body, "frequency_penalty", _frequency_penalty); addMember(body, "temperature", _temperature); addMember(body, "top_p", _top_p); if (_stream) { addMember(body, "stream", true); Value8 stream_options; addMember(stream_options, "include_usage", _stream_options_include_usage); addMember(body, "stream_options", stream_options); } std::string json(write2Json(body, false)); return json; } void SHAI::Request::addMessage(const roleType &rt, LPCSTR content) { message msg; msg.content = content; msg.role = rt; _messages.push_back(msg); } void SHAI::Request::addUserMessage(LPCSTR content) { addMessage(SHAI::Request::RT_USER, content); } void SHAI::Request::addSystemMessage(LPCSTR content) { addMessage(SHAI::Request::RT_SYSTEM, content); } void SHAI::Request::addAssistantMessage(LPCSTR content) { addMessage(SHAI::Request::RT_ASSISTANT, content); } void SHAI::Request::clearMessages() { _messages.clear(); } void SHAI::Request::addHeaders(LPCSTR key, LPCSTR value) { _headers.insert(std::make_pair(key, value)); } std::string SHAI::Request::apiEndPoint() const { return _apiEndpoint; } void SHAI::Request::popbackMessage() { _messages.pop_back(); } std::string SHAI::Request::getRole(const roleType &rt) const { switch (rt) { case RT_SYSTEM: return "system"; case RT_USER: return "user"; case RT_ASSISTANT: return "assistant"; case RT_TOOL: return "tool"; } return "user"; } SHAI::Request::Request(LPCSTR apikey , LPCSTR model , LPCSTR url , const bool &stream) : _apiKey(apikey) , _model(model) , _apiEndpoint(url) , _stream(stream) , _stream_options_include_usage(false) { if (_stream) _stream_options_include_usage = true; } SHAI::Request::Request() : _stream(true) , _stream_options_include_usage(false) , _max_tokens(4096) , _top_p(.7f) , _top_k(50) , _frequency_penalty(.0f) , _presence_penalty(.0f) , _temperature(.6f) { if (_stream) _stream_options_include_usage = true; } int SHAI::Request::send(Response &response) const { // 1. 准备请求头 std::string authHeader = std::string("Authorization: Bearer ") + _apiKey; const char *headers[2] = { "Content-Type: application/json", authHeader.c_str() }; // 2. 准备请求体 (依然用您的 RapidJSON 生成) std::string sBody(requestBody()); // 3. 一句话调用 DLL!把网络的脏活累活全丢过去 int resultCode = mg_curl_httpPost( apiEndPoint().c_str(), // URL headers, // 请求头数组 2, // 请求头数量 sBody.c_str(), // JSON body request_WriteCallback, // 您 ARX 里的解析函数 &response // 透传指针 ); return resultCode; } #pragma endregion #pragma region Response void SHAI::Response::clear() { _reasoning_content.clear(); _content.clear(); _all_reasoning_content.clear(); _all_content.clear(); _bUsageTokens = false; _prompt_tokens = _completion_tokens = _total_tokens = 0; _allOriginalResponse.clear(); _bRecordContext = true; _finish_reason.clear(); } void SHAI::Response::parseChoices(CValue8 &v) { CMIt8 itChoices = v.FindMember("choices"); if (itChoices == v.MemberEnd()) return; if (!itChoices->value.IsArray()) return; // 目前只处理第一个 CArr8 ar = itChoices->value.GetArray(); for (CVIt8 itAr = ar.Begin(); itAr != ar.End(); ++itAr) { CMIt8 itItem = itAr->FindMember("finish_reason"); if (itItem != itAr->MemberEnd()) { if (!itItem->value.IsNull()) _finish_reason = itItem->value.GetString(); } itItem = itAr->FindMember("delta"); if (itItem != itAr->MemberEnd()) { CValue8 &vDelta = itItem->value; CMIt8 itContent = vDelta.FindMember("content"); if (itContent != vDelta.MemberEnd() && !itContent->value.IsNull()) { _content = itContent->value.GetString(); _all_content += _content; } CMIt8 itReasoningContent = vDelta.FindMember("reasoning_content"); if (itReasoningContent != vDelta.MemberEnd() && !itReasoningContent->value.IsNull()) { _reasoning_content = itReasoningContent->value.GetString(); _all_reasoning_content += _reasoning_content; } CMIt8 itRole = vDelta.FindMember("role"); if (itRole != vDelta.MemberEnd()) { //m["role"] = itRole->value.GetString(); } } } } SHAI::Response::Response() : _bUsageTokens(false) , _prompt_tokens(0) , _completion_tokens(0) , _total_tokens(0) , _bShowThinkingContent(true) , _bRecordContext(true) { } void SHAI::Response::readResponse(const std::string &s) { static bool bStartThinking(false), bStartAnswer(false); appendOriginalResponse(s); if (s.find("[DONE]") != std::string::npos) { bStartThinking = false; bStartAnswer = false; return; } _content = _reasoning_content = ""; Document8 doc; doc.Parse(s); if (doc.HasParseError()) return; if (!doc.IsArray() && !doc.IsObject()) return; parseChoices(doc); CMIt8 itUsage = doc.FindMember("usage"); if (itUsage != doc.MemberEnd() && !itUsage->value.IsNull() && _finish_reason == "stop" && !_bUsageTokens) { _bUsageTokens = true; CValue8 &vUsage(itUsage->value); CMIt8 itPromptTokens = vUsage.FindMember("prompt_tokens"); if (itPromptTokens != vUsage.MemberEnd()) _prompt_tokens = itPromptTokens->value.GetInt(); CMIt8 itCompletionTokens = vUsage.FindMember("completion_tokens"); if (itCompletionTokens != vUsage.MemberEnd()) _completion_tokens = itCompletionTokens->value.GetInt(); CMIt8 itTotalTokens = vUsage.FindMember("total_tokens"); if (itTotalTokens != vUsage.MemberEnd()) _total_tokens = itTotalTokens->value.GetInt(); CString str; str.Format(_T("\n⌈输入tokens:%d | 输出tokens:%d⌋\n") , promptTokens() , completionTokens() ); g_aiPaletteSet::instance().sendTextToAnswer(str); } if (!content().empty()) { if (!bStartAnswer) { bStartAnswer = true; g_aiPaletteSet::instance().sendTextToAnswer(_T("\r\n\r\n回答:\r\n")); } g_aiPaletteSet::instance().sendTextToAnswer(AIJSON::utf2w(content()).c_str()); } if (!reasoningContent().empty() && showThinkingContent()) { if (!bStartThinking) { bStartThinking = true; g_aiPaletteSet::instance().sendTextToAnswer(_T("\r\n\r\n思考:\r\n")); } g_aiPaletteSet::instance().sendTextToAnswer(AIJSON::utf2w(reasoningContent()).c_str()); } } #pragma endregion #pragma region AICloud // 设置获取 当前模型 void SHAI::AICloud::setCurrentModel(const int &idx) { _current_model.clear(); if (idx >= 0 && idx < _models.size()) _current_model = _models[idx]; } void SHAI::AICloud::parseModelsJson(LPCSTR ss, std::vector &models) { Document8 doc; doc.Parse(ss); if (doc.HasParseError()) return; if (!doc.IsArray() && !doc.IsObject()) return; CMIt8 itData = doc.FindMember("data"); if (itData == doc.MemberEnd()) return; if (!itData->value.IsArray()) return; const CArr8 &arData = itData->value.GetArray(); CVIt8 itArr = arData.Begin(); CMIt8 itModel; for (; itArr != arData.End(); ++itArr) { itModel = itArr->FindMember("id"); if (itModel == itArr->MemberEnd()) continue; models.push_back(itModel->value.GetString()); } } SHAI::AICloud *SHAI::AICloud::createCloud(LPCSTR name) { if (AIJSON::w2utf(_T("阿里云百炼")).compare(name) == 0) { return new AliyunBailian(); } else if (AIJSON::w2utf(_T("硅基流动")).compare(name) == 0) { return new SiliconFlow(); } else if (std::string(name) == "DeepSeek") { return new DeepSeek(); } else if (std::string(name) == "Hyperbolic") { return new Hyperbolic(); } else { return new CustomCloud(name, "输入API key", "输入API网址"); } return NULL; } bool SHAI::AICloud::IsCustomCloud(AICloud *p) { std::string name = p->cloudName(); if (AIJSON::w2utf(_T("阿里云百炼")).compare(name) == 0) { return false; } else if (AIJSON::w2utf(_T("硅基流动")).compare(name) == 0) { return false; } else if (std::string(name) == "DeepSeek") { return false; } else if (std::string(name) == "Hyperbolic") { return false; } return true; } void SHAI::AliyunBailian::getModels() { _models.clear(); _models.push_back("deepseek-r1"); _models.push_back("deepseek-v3"); } void SHAI::SiliconFlow::getModels() { _models.clear(); // 1. 准备请求头 std::string authHeader = std::string("Authorization: Bearer ") + _apiKey; const char* headers[1] = { authHeader.c_str() }; // 2. 准备接收返回数据的 string std::string responseString; // 3. 直接调用胶水层的 GET 方法 (假设您的回调函数已经重写过) int resultCode = mg_curl_httpGet( "https://api.siliconflow.cn/v1/models", headers, 1, get_model_list_WriteCallback, // 专门用来接收普通字符串的回调 &responseString ); if (resultCode == 0) // 0 对应 CURLE_OK { AICloud::parseModelsJson(responseString.c_str(), _models); if (_models.empty()) AfxMessageBox(_T("获取模型列表失败。")); } else { AfxMessageBox(_T("网络请求失败或超时。")); } } void SHAI::DeepSeek::getModels() { _models.clear(); // 1. 准备请求头 std::string authHeader = std::string("Authorization: Bearer ") + _apiKey; const char* headers[2] = { "Accept: application/json", authHeader.c_str() }; // 2. 准备接收返回数据的 string std::string responseString; // 3. 直接调用胶水层的 GET 方法 (假设您的回调函数已经重写过) int resultCode = mg_curl_httpGet( "https://api.deepseek.com/models", headers, 2, get_model_list_WriteCallback, // 专门用来接收普通字符串的回调 &responseString ); if (resultCode == 0) // 0 对应 CURLE_OK { AICloud::parseModelsJson(responseString.c_str(), _models); if (_models.empty()) AfxMessageBox(_T("获取模型列表失败。")); } else { AfxMessageBox(_T("网络请求失败或超时。")); } } void SHAI::Hyperbolic::getModels() { _models.clear(); _models.push_back("deepseek-ai/DeepSeek-R1-Zero"); _models.push_back("deepseek-ai/DeepSeek-R1"); _models.push_back("deepseek-ai/DeepSeek-V3"); _models.push_back("meta-llama/Llama-3.3-70B-Instruct"); } #pragma endregion