Files
didactyl/src/tools/tool_model.c
T

326 lines
9.9 KiB
C

#define _POSIX_C_SOURCE 200809L
#include "tools_internal.h"
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include "cjson/cJSON.h"
#include "../config.h"
#include "../llm.h"
static char* json_error_local(const char* msg) {
cJSON* root = cJSON_CreateObject();
if (!root) return NULL;
cJSON_AddBoolToObject(root, "success", 0);
cJSON_AddStringToObject(root, "error", msg ? msg : "unknown error");
char* out = cJSON_PrintUnformatted(root);
cJSON_Delete(root);
return out;
}
static cJSON* parse_args_local(const char* args_json) {
const char* raw = args_json ? args_json : "{}";
cJSON* args = cJSON_Parse(raw);
if (!args || !cJSON_IsObject(args)) {
cJSON_Delete(args);
return NULL;
}
return args;
}
static int json_object_set_string(cJSON* obj, const char* key, const char* value) {
if (!obj || !key || !value) return -1;
cJSON_DeleteItemFromObjectCaseSensitive(obj, key);
cJSON* v = cJSON_CreateString(value);
if (!v) return -1;
cJSON_AddItemToObject(obj, key, v);
return 0;
}
static char* read_entire_file_local(const char* file_path, size_t* out_bytes) {
if (!file_path) return NULL;
FILE* fp = fopen(file_path, "rb");
if (!fp) return NULL;
if (fseek(fp, 0, SEEK_END) != 0) {
fclose(fp);
return NULL;
}
long len = ftell(fp);
if (len < 0) {
fclose(fp);
return NULL;
}
if (fseek(fp, 0, SEEK_SET) != 0) {
fclose(fp);
return NULL;
}
char* buf = (char*)malloc((size_t)len + 1U);
if (!buf) {
fclose(fp);
return NULL;
}
size_t n = fread(buf, 1, (size_t)len, fp);
fclose(fp);
if (n != (size_t)len) {
free(buf);
return NULL;
}
buf[len] = '\0';
if (out_bytes) *out_bytes = (size_t)len;
return buf;
}
static int persist_llm_config(tools_context_t* ctx, const llm_config_t* cfg) {
if (!ctx || !ctx->cfg || !cfg) return -1;
if (ctx->cfg->config_path[0] == '\0') return -1;
size_t src_len = 0;
char* raw = read_entire_file_local(ctx->cfg->config_path, &src_len);
if (!raw) return -1;
char* src = jsonc_strip_comments(raw, src_len);
free(raw);
if (!src) return -1;
cJSON* root = cJSON_ParseWithLength(src, strlen(src));
free(src);
if (!root || !cJSON_IsObject(root)) {
cJSON_Delete(root);
return -1;
}
cJSON* llm = cJSON_GetObjectItemCaseSensitive(root, "llm");
if (!llm || !cJSON_IsObject(llm)) {
cJSON_DeleteItemFromObjectCaseSensitive(root, "llm");
llm = cJSON_CreateObject();
if (!llm) {
cJSON_Delete(root);
return -1;
}
cJSON_AddItemToObject(root, "llm", llm);
}
if (json_object_set_string(llm, "provider", cfg->provider) != 0 ||
json_object_set_string(llm, "api_key", cfg->api_key) != 0 ||
json_object_set_string(llm, "model", cfg->model) != 0 ||
json_object_set_string(llm, "base_url", cfg->base_url) != 0) {
cJSON_Delete(root);
return -1;
}
cJSON_DeleteItemFromObjectCaseSensitive(llm, "max_tokens");
cJSON_AddNumberToObject(llm, "max_tokens", cfg->max_tokens);
cJSON_DeleteItemFromObjectCaseSensitive(llm, "temperature");
cJSON_AddNumberToObject(llm, "temperature", cfg->temperature);
char* out = cJSON_Print(root);
cJSON_Delete(root);
if (!out) return -1;
FILE* fp = fopen(ctx->cfg->config_path, "wb");
if (!fp) {
free(out);
return -1;
}
size_t out_len = strlen(out);
size_t n = fwrite(out, 1, out_len, fp);
fclose(fp);
free(out);
return n == out_len ? 0 : -1;
}
static int assign_string_field(cJSON* item, char* dst, size_t dst_size, int* changed) {
if (!item) return 0;
if (!cJSON_IsString(item) || !item->valuestring) return -1;
size_t n = strlen(item->valuestring);
if (n >= dst_size) return -1;
memcpy(dst, item->valuestring, n + 1U);
if (changed) *changed = 1;
return 0;
}
char* execute_model_get(const char* args_json) {
cJSON* args = parse_args_local(args_json);
if (!args) return json_error_local("invalid arguments JSON");
cJSON_Delete(args);
llm_config_t cfg;
if (llm_get_config(&cfg) != 0) {
return json_error_local("llm runtime unavailable");
}
cJSON* out = cJSON_CreateObject();
if (!out) return NULL;
cJSON_AddBoolToObject(out, "success", 1);
cJSON_AddStringToObject(out, "provider", cfg.provider);
cJSON_AddStringToObject(out, "model", cfg.model);
cJSON_AddStringToObject(out, "base_url", cfg.base_url);
cJSON_AddNumberToObject(out, "max_tokens", cfg.max_tokens);
cJSON_AddNumberToObject(out, "temperature", cfg.temperature);
char* json = cJSON_PrintUnformatted(out);
cJSON_Delete(out);
return json;
}
char* execute_model_set(tools_context_t* ctx, const char* args_json) {
if (!ctx || !ctx->cfg) return json_error_local("tool context unavailable");
cJSON* args = parse_args_local(args_json);
if (!args) return json_error_local("invalid arguments JSON");
llm_config_t cfg;
if (llm_get_config(&cfg) != 0) {
cJSON_Delete(args);
return json_error_local("llm runtime unavailable");
}
int changed = 0;
cJSON* provider = cJSON_GetObjectItemCaseSensitive(args, "provider");
cJSON* api_key = cJSON_GetObjectItemCaseSensitive(args, "api_key");
cJSON* model = cJSON_GetObjectItemCaseSensitive(args, "model");
cJSON* base_url = cJSON_GetObjectItemCaseSensitive(args, "base_url");
cJSON* max_tokens = cJSON_GetObjectItemCaseSensitive(args, "max_tokens");
cJSON* temperature = cJSON_GetObjectItemCaseSensitive(args, "temperature");
if (assign_string_field(provider, cfg.provider, sizeof(cfg.provider), &changed) != 0 ||
assign_string_field(api_key, cfg.api_key, sizeof(cfg.api_key), &changed) != 0 ||
assign_string_field(model, cfg.model, sizeof(cfg.model), &changed) != 0 ||
assign_string_field(base_url, cfg.base_url, sizeof(cfg.base_url), &changed) != 0) {
cJSON_Delete(args);
return json_error_local("model_set string field invalid or too long");
}
if (max_tokens) {
if (!cJSON_IsNumber(max_tokens)) {
cJSON_Delete(args);
return json_error_local("model_set max_tokens must be a number");
}
cfg.max_tokens = (int)max_tokens->valuedouble;
changed = 1;
}
if (temperature) {
if (!cJSON_IsNumber(temperature)) {
cJSON_Delete(args);
return json_error_local("model_set temperature must be a number");
}
cfg.temperature = temperature->valuedouble;
changed = 1;
}
cJSON_Delete(args);
if (!changed) {
return json_error_local("model_set requires at least one field to update");
}
if (llm_set_config(&cfg) != 0) {
return json_error_local("failed to update runtime llm config");
}
ctx->cfg->llm = cfg;
if (persist_llm_config(ctx, &cfg) != 0) {
return json_error_local("failed to persist llm config to config file");
}
cJSON* out = cJSON_CreateObject();
if (!out) return NULL;
cJSON_AddBoolToObject(out, "success", 1);
cJSON_AddStringToObject(out, "provider", cfg.provider);
cJSON_AddStringToObject(out, "model", cfg.model);
cJSON_AddStringToObject(out, "base_url", cfg.base_url);
cJSON_AddNumberToObject(out, "max_tokens", cfg.max_tokens);
cJSON_AddNumberToObject(out, "temperature", cfg.temperature);
char* json = cJSON_PrintUnformatted(out);
cJSON_Delete(out);
return json;
}
static void append_model_id(cJSON* ids, cJSON* item) {
if (!ids || !item) return;
if (cJSON_IsString(item) && item->valuestring) {
cJSON_AddItemToArray(ids, cJSON_CreateString(item->valuestring));
return;
}
if (cJSON_IsObject(item)) {
cJSON* id = cJSON_GetObjectItemCaseSensitive(item, "id");
if (id && cJSON_IsString(id) && id->valuestring) {
cJSON_AddItemToArray(ids, cJSON_CreateString(id->valuestring));
}
}
}
char* execute_model_list(const char* args_json) {
cJSON* args = parse_args_local(args_json);
if (!args) return json_error_local("invalid arguments JSON");
cJSON* base_url = cJSON_GetObjectItemCaseSensitive(args, "base_url");
if (base_url && (!cJSON_IsString(base_url) || !base_url->valuestring)) {
cJSON_Delete(args);
return json_error_local("model_list base_url must be a string");
}
const char* base_url_override = (base_url && base_url->valuestring && base_url->valuestring[0] != '\0')
? base_url->valuestring
: NULL;
char* raw = llm_list_models_json(base_url_override);
cJSON_Delete(args);
if (!raw) {
return json_error_local("model_list request failed");
}
cJSON* root = cJSON_Parse(raw);
free(raw);
if (!root) {
cJSON_Delete(root);
return json_error_local("model_list returned invalid JSON");
}
cJSON* ids = cJSON_CreateArray();
cJSON* out = cJSON_CreateObject();
if (!ids || !out) {
cJSON_Delete(ids);
cJSON_Delete(out);
cJSON_Delete(root);
return NULL;
}
if (cJSON_IsArray(root)) {
int n = cJSON_GetArraySize(root);
for (int i = 0; i < n; i++) {
append_model_id(ids, cJSON_GetArrayItem(root, i));
}
} else if (cJSON_IsObject(root)) {
cJSON* data = cJSON_GetObjectItemCaseSensitive(root, "data");
if (data && cJSON_IsArray(data)) {
int n = cJSON_GetArraySize(data);
for (int i = 0; i < n; i++) {
append_model_id(ids, cJSON_GetArrayItem(data, i));
}
}
}
cJSON_AddBoolToObject(out, "success", 1);
cJSON_AddNumberToObject(out, "count", cJSON_GetArraySize(ids));
cJSON_AddItemToObject(out, "models", ids);
char* json = cJSON_PrintUnformatted(out);
cJSON_Delete(out);
cJSON_Delete(root);
return json;
}