Files
envi-code/SourceCode/Code2026/CoreCommand/SH_AI.cpp
T
2026-09-28 15:28:50 +08:00

659 lines
15 KiB
C++
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#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<std::string> &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<std::string> 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<std::string> &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