659 lines
15 KiB
C++
659 lines
15 KiB
C++
#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
|
||
|