feat(ai): add model caching, settings, and commands
Some checks failed
CI Code / Check spelling (pull_request) Successful in 23s
CI Code / Check coding style (pull_request) Successful in 35s
CI Code / Code Coverage (pull_request) Successful in 2m28s
CI Code / Linux (debian) (pull_request) Failing after 6m18s
CI Code / Linux (arch) (pull_request) Failing after 8m46s
CI Code / Linux (ubuntu) (pull_request) Failing after 11m4s

- Introduce model caching with persistence to preferences
- Add provider default model and custom settings management
- Implement `/ai switch`, `/ai models`, and improve `/ai start`
- Add model name autocomplete for chat commands
- Update command definitions and help text
- Add unit tests for new functionality
This commit is contained in:
2026-05-09 13:15:58 +00:00
parent 731b55fa19
commit acae543057
11 changed files with 1255 additions and 98 deletions

View File

@@ -32,6 +32,12 @@ static GHashTable* providers = NULL;
static GHashTable* provider_keys = NULL;
static Autocomplete providers_ac = NULL;
/* ========================================================================
* Forward declarations
* ======================================================================== */
static void _ai_load_models_for_provider(AIProvider* provider);
static void _ai_save_models_for_provider(AIProvider* provider);
/* ========================================================================
* Curl helpers
* ======================================================================== */
@@ -183,7 +189,10 @@ ai_provider_new(const gchar* name, const gchar* api_url, const gchar* org_id)
provider->api_url = g_strdup(api_url ? api_url : "");
provider->org_id = g_strdup(org_id);
provider->project_id = NULL;
provider->default_model = NULL;
provider->settings = g_hash_table_new_full(g_str_hash, g_str_equal, g_free, g_free);
provider->models = NULL;
provider->models_fresh = FALSE;
provider->ref_count = 1;
return provider;
}
@@ -211,6 +220,8 @@ ai_provider_unref(AIProvider* provider)
g_free(provider->api_url);
g_free(provider->org_id);
g_free(provider->project_id);
g_free(provider->default_model);
g_hash_table_destroy(provider->settings);
GList* curr = provider->models;
while (curr) {
@@ -271,6 +282,16 @@ ai_client_init(void)
/* Load saved API keys from config */
ai_load_keys();
/* Load cached models for all providers */
GList* provider_list = ai_list_providers();
GList* curr = provider_list;
while (curr) {
AIProvider* provider = (AIProvider*)curr->data;
_ai_load_models_for_provider(provider);
curr = g_list_next(curr);
}
g_list_free(provider_list);
log_info("AI client initialized with default providers: openai, perplexity");
}
@@ -387,6 +408,37 @@ ai_providers_find(const char* const search_str, gboolean previous, void* context
return autocomplete_complete(providers_ac, effective_search, FALSE, previous);
}
/* Stateful model name finder for autocomplete.
* Uses the current session's provider models for completion. */
gchar*
ai_models_find(const char* const search_str, gboolean previous, void* context)
{
ProfAiWin* aiwin = (ProfAiWin*)context;
if (!aiwin || !aiwin->session) {
return NULL;
}
AIProvider* provider = aiwin->session->provider;
if (!provider || !provider->models) {
return NULL;
}
/* Create a temporary autocomplete for models */
Autocomplete models_ac = autocomplete_new();
GList* curr = provider->models;
while (curr) {
autocomplete_add(models_ac, curr->data);
curr = g_list_next(curr);
}
/* NULL search_str is treated as empty string for cycling */
const char* effective_search = (search_str != NULL) ? search_str : "";
gchar* result = autocomplete_complete(models_ac, effective_search, FALSE, previous);
autocomplete_free(models_ac);
return result;
}
gchar*
ai_get_provider_key(const gchar* provider_name)
{
@@ -419,6 +471,456 @@ ai_set_provider_key(const gchar* provider_name, const gchar* api_key)
}
}
/* ========================================================================
* Provider default model and settings
* ======================================================================== */
void
ai_set_provider_default_model(const gchar* provider_name, const gchar* model)
{
if (!provider_name || !model)
return;
AIProvider* provider = ai_get_provider(provider_name);
if (!provider) {
log_warning("Provider '%s' not found for default model setting", provider_name);
return;
}
g_free(provider->default_model);
provider->default_model = g_strdup(model);
log_info("Default model for provider '%s' set to: %s", provider_name, model);
}
const gchar*
ai_get_provider_default_model(const gchar* provider_name)
{
if (!provider_name)
return NULL;
AIProvider* provider = ai_get_provider(provider_name);
if (!provider)
return NULL;
return provider->default_model;
}
void
ai_set_provider_setting(const gchar* provider_name, const gchar* setting, const gchar* value)
{
if (!provider_name || !setting)
return;
AIProvider* provider = ai_get_provider(provider_name);
if (!provider) {
log_warning("Provider '%s' not found for setting '%s'", provider_name, setting);
return;
}
if (!provider->settings) {
provider->settings = g_hash_table_new_full(g_str_hash, g_str_equal, g_free, g_free);
}
if (value) {
g_hash_table_insert(provider->settings, g_strdup(setting), g_strdup(value));
log_info("Setting '%s' for provider '%s' set to: %s", setting, provider_name, value);
} else {
g_hash_table_remove(provider->settings, setting);
log_info("Setting '%s' removed for provider '%s'", setting, provider_name);
}
}
gchar*
ai_get_provider_setting(const gchar* provider_name, const gchar* setting)
{
if (!provider_name || !setting)
return NULL;
AIProvider* provider = ai_get_provider(provider_name);
if (!provider || !provider->settings)
return NULL;
gchar* value = g_hash_table_lookup(provider->settings, setting);
return g_strdup(value);
}
/* ========================================================================
* Generic HTTP request handling
* ======================================================================== */
typedef void (*ai_response_handler_t)(AIProvider* provider, const gchar* response, gpointer user_data);
typedef struct
{
gchar* provider_name;
gpointer user_data;
gchar* request_url;
ai_response_handler_t handler;
} ai_request_t;
/**
* Build curl headers with auth and optional org header.
* Returns headers list (caller must free with curl_slist_free_all).
*/
static struct curl_slist*
_build_curl_headers(AIProvider* provider, const gchar* api_key)
{
struct curl_slist* headers = NULL;
auto_gchar gchar* auth_header = g_strdup_printf("Authorization: Bearer %s", api_key);
headers = curl_slist_append(headers, "Content-Type: application/json");
headers = curl_slist_append(headers, auth_header);
if (provider && provider->org_id && strlen(provider->org_id) > 0) {
auto_gchar gchar* org_header = g_strdup_printf("OpenAI-Organization: %s", provider->org_id);
headers = curl_slist_append(headers, org_header);
}
return headers;
}
/**
* Execute a curl request and call handler on success.
* Returns CURLcode.
*/
static CURLcode
_curl_exec_and_handle(CURL* curl, const gchar* url, struct curl_slist* headers,
struct curl_response_t* response, ai_response_handler_t handler,
AIProvider* provider, gpointer user_data)
{
curl_easy_setopt(curl, CURLOPT_URL, url);
curl_easy_setopt(curl, CURLOPT_HTTPHEADER, headers);
curl_easy_setopt(curl, CURLOPT_WRITEFUNCTION, _write_callback);
curl_easy_setopt(curl, CURLOPT_WRITEDATA, response);
curl_easy_setopt(curl, CURLOPT_TIMEOUT, 30L);
CURLcode res = curl_easy_perform(curl);
long http_code = 0;
curl_easy_getinfo(curl, CURLINFO_RESPONSE_CODE, &http_code);
if (res != CURLE_OK) {
auto_gchar gchar* error_msg = g_strdup(curl_easy_strerror(res));
log_error("Request failed for '%s': %s", provider->name, error_msg);
_aiwin_display_error(user_data, error_msg);
} else if (http_code >= 400) {
auto_gchar gchar* error_msg = g_strdup_printf("HTTP %ld: %s", http_code,
response->data ? response->data : "Unknown error");
log_error("Request HTTP error for '%s': %s", provider->name, error_msg);
_aiwin_display_error(user_data, error_msg);
} else {
handler(provider, response->data, user_data);
}
return res;
}
static gpointer
_ai_generic_request_thread(gpointer data)
{
ai_request_t* req = (ai_request_t*)data;
AIProvider* provider = ai_get_provider(req->provider_name);
if (!provider) {
log_error("Provider '%s' not found for request", req->provider_name);
g_free(req->provider_name);
g_free(req->request_url);
g_free(req);
return NULL;
}
CURL* curl = curl_easy_init();
if (!curl) {
log_error("Failed to initialize curl for %s", req->provider_name);
g_free(req->provider_name);
g_free(req->request_url);
g_free(req);
return NULL;
}
auto_gchar gchar* api_key = ai_get_provider_key(req->provider_name);
if (!api_key || strlen(api_key) == 0) {
auto_gchar gchar* error_msg = g_strdup_printf("No API key set for provider '%s'.", req->provider_name);
_aiwin_display_error(req->user_data, error_msg);
curl_easy_cleanup(curl);
g_free(req->provider_name);
g_free(req->request_url);
g_free(req);
return NULL;
}
struct curl_slist* headers = _build_curl_headers(provider, api_key);
struct curl_response_t response;
response.data = g_new0(gchar, 1);
response.size = 0;
_curl_exec_and_handle(curl, req->request_url, headers, &response, req->handler, provider, req->user_data);
curl_slist_free_all(headers);
curl_easy_cleanup(curl);
g_free(response.data);
g_free(req->provider_name);
g_free(req->request_url);
g_free(req);
return NULL;
}
/* ========================================================================
* Model fetching
* ======================================================================== */
gboolean
ai_models_are_fresh(const gchar* provider_name)
{
if (!provider_name)
return FALSE;
AIProvider* provider = ai_get_provider(provider_name);
if (!provider)
return FALSE;
return provider->models_fresh;
}
static gboolean
_model_exists(AIProvider* provider, const gchar* model)
{
GList* curr = provider->models;
while (curr) {
if (g_strcmp0(curr->data, model) == 0) {
return TRUE;
}
curr = g_list_next(curr);
}
return FALSE;
}
static void
_add_model_if_new(AIProvider* provider, const gchar* model)
{
if (!_model_exists(provider, model)) {
provider->models = g_list_append(provider->models, g_strdup(model));
}
}
/* ========================================================================
* Model cache persistence
* ======================================================================== */
/**
* Load cached models for a provider from preferences.
*/
static void
_ai_load_models_for_provider(AIProvider* provider)
{
if (!provider)
return;
GList* models = prefs_ai_get_models(provider->name);
if (!models) {
log_info("No cached models for provider '%s'", provider->name);
return;
}
/* Add loaded models to provider */
GList* curr = models;
while (curr) {
_add_model_if_new(provider, curr->data);
curr = g_list_next(curr);
}
prefs_free_ai_models(models);
if (provider->models) {
provider->models_fresh = TRUE;
log_info("Loaded %d cached models for provider '%s'", g_list_length(provider->models), provider->name);
}
}
/**
* Save models for a provider to preferences.
*/
static void
_ai_save_models_for_provider(AIProvider* provider)
{
if (!provider || !provider->models)
return;
/* Build array of model strings */
gint count = g_list_length(provider->models);
gchar** models = g_new0(gchar*, count);
GList* curr = provider->models;
gint i = 0;
while (curr) {
models[i++] = curr->data;
curr = g_list_next(curr);
}
prefs_ai_set_models(provider->name, (const gchar* const*)models, count);
g_free(models);
log_debug("Saved %d models for provider '%s' to prefs", count, provider->name);
}
static const gchar*
_find_json_field(const gchar* json, const gchar* field)
{
auto_gchar gchar* key = g_strdup_printf("\"%s\":\"", field);
const gchar* pos = strstr(json, key);
if (!pos) {
return NULL;
}
return pos + strlen(key);
}
static void
_parse_openai_models(AIProvider* provider, const gchar* json)
{
/* Format: {"data": [{"id": "model1"}, {"id": "model2"}, ...]} */
const gchar* data_start = _find_json_field(json, "data");
if (!data_start) {
return;
}
const gchar* obj_start = strstr(data_start, "{\"id\":");
while (obj_start) {
const gchar* id_start = _find_json_field(obj_start, "id");
if (!id_start) {
break;
}
const gchar* id_end = strchr(id_start, '"');
if (!id_end) {
break;
}
auto_gchar gchar* model = g_strndup(id_start, id_end - id_start);
_add_model_if_new(provider, model);
obj_start = strstr(id_end, "{\"id\":");
if (obj_start) {
obj_start++;
}
}
}
static void
_parse_array_models(AIProvider* provider, const gchar* json)
{
/* Format: ["model1", "model2", ...] */
const gchar* arr_start = strchr(json, '[');
if (!arr_start) {
return;
}
const gchar* model_start = strchr(arr_start, '"');
while (model_start) {
const gchar* model_end = strchr(model_start + 1, '"');
if (!model_end) {
break;
}
auto_gchar gchar* model = g_strndup(model_start + 1, model_end - model_start - 1);
_add_model_if_new(provider, model);
model_start = strchr(model_end + 1, '"');
}
}
static void
_parse_and_cache_models(AIProvider* provider, const gchar* response_json)
{
if (!provider || !response_json) {
return;
}
/* Try OpenAI format first: {"data": [{"id": "model1"}, ...]} */
_parse_openai_models(provider, response_json);
/* If no models found, try simple array format: ["model1", ...] */
if (!provider->models) {
_parse_array_models(provider, response_json);
}
provider->models_fresh = TRUE;
log_info("Cached %d models for provider '%s'", g_list_length(provider->models), provider->name);
/* Persist models to preferences */
_ai_save_models_for_provider(provider);
}
static void
_models_response_handler(AIProvider* provider, const gchar* response, gpointer user_data)
{
_parse_and_cache_models(provider, response);
auto_gchar gchar* msg = g_strdup_printf("Cached %d models for provider '%s'. Use '/ai start <provider> <model>' to start a chat.",
g_list_length(provider->models), provider->name);
_aiwin_display_response(user_data, msg);
}
static gpointer
_models_fetch_thread(gpointer data)
{
ai_request_t* req = (ai_request_t*)data;
AIProvider* provider = ai_get_provider(req->provider_name);
if (!provider) {
g_free(req->provider_name);
g_free(req);
return NULL;
}
/* Build URL for listing models */
auto_gchar gchar* base_url = g_strdup_printf("%s%s", provider->api_url,
g_str_has_suffix(provider->api_url, "/") ? "" : "/");
auto_gchar gchar* request_url = g_strdup_printf("%sv1/models", base_url);
req->request_url = g_strdup(request_url);
req->handler = _models_response_handler;
return _ai_generic_request_thread(req);
}
gboolean
ai_fetch_models(const gchar* provider_name, gpointer user_data)
{
if (!provider_name)
return FALSE;
AIProvider* provider = ai_get_provider(provider_name);
if (!provider) {
log_error("Provider '%s' not found for model fetch", provider_name);
return FALSE;
}
/* If models are fresh, just display them */
if (provider->models_fresh && provider->models) {
cons_show("Cached %d models for provider '%s':", g_list_length(provider->models), provider_name);
GList* curr = provider->models;
while (curr) {
cons_show(" %s", (gchar*)curr->data);
curr = g_list_next(curr);
}
cons_show("");
cons_show("Use '/ai models %s --refresh' to fetch fresh models from the API.", provider_name);
return TRUE;
}
/* Fetch models in a thread */
ai_request_t* req = g_new0(ai_request_t, 1);
req->provider_name = g_strdup(provider_name);
req->user_data = user_data;
GThread* thread = g_thread_new("ai_models_fetch", _models_fetch_thread, req);
if (!thread) {
g_free(req->provider_name);
g_free(req);
log_error("Failed to create thread for model fetch");
return FALSE;
}
g_thread_unref(thread);
cons_show("Fetching models for provider '%s'...", provider_name);
return TRUE;
}
/* ========================================================================
* Session Management
* ======================================================================== */
@@ -708,34 +1210,23 @@ _ai_request_thread(gpointer data)
auto_gchar gchar* json_payload = _build_json_payload(session, prompt);
log_debug("[AI-THREAD] JSON payload: %s", json_payload);
/* Set up headers */
struct curl_slist* headers = NULL;
auto_gchar gchar* auth_header = g_strdup_printf("Authorization: Bearer %s", session->api_key);
headers = curl_slist_append(headers, "Content-Type: application/json");
headers = curl_slist_append(headers, auth_header);
/* Add organization header if configured */
if (session->provider && session->provider->org_id && strlen(session->provider->org_id) > 0) {
auto_gchar gchar* org_header = g_strdup_printf("OpenAI-Organization: %s", session->provider->org_id);
headers = curl_slist_append(headers, org_header);
}
/* Build headers using shared helper */
struct curl_slist* headers = _build_curl_headers(session->provider, session->api_key);
/* Response buffer */
struct curl_response_t response;
response.data = g_new0(gchar, 1);
response.size = 0;
/* Configure request */
/* Build request URL */
const gchar* api_url = session->provider ? session->provider->api_url : DEFAULT_OPENAI_URL;
auto_gchar gchar* request_url = g_strdup_printf("%s%sv1/responses", api_url, g_str_has_suffix(api_url, "/") ? "" : "/");
log_debug("[AI-THREAD] API URL: %s", api_url);
log_debug("[AI-THREAD] API Request URL: %s", request_url);
log_debug("[AI-THREAD] Model: %s", session->model);
curl_easy_setopt(curl, CURLOPT_URL, request_url);
curl_easy_setopt(curl, CURLOPT_HTTPHEADER, headers);
/* Execute request with POST fields */
curl_easy_setopt(curl, CURLOPT_POSTFIELDS, json_payload);
curl_easy_setopt(curl, CURLOPT_WRITEFUNCTION, _write_callback);
curl_easy_setopt(curl, CURLOPT_WRITEDATA, &response);
curl_easy_setopt(curl, CURLOPT_TIMEOUT, 60L);
CURLcode res = curl_easy_perform(curl);

View File

@@ -17,12 +17,15 @@ typedef struct ai_message_t
*/
typedef struct ai_provider_t
{
gchar* name; /* Provider name (e.g., "openai", "perplexity") */
gchar* api_url; /* API endpoint URL */
gchar* org_id; /* Optional organization ID */
gchar* project_id; /* Optional project ID (for some providers) */
GList* models; /* List of available models (gchar*) */
gint32 ref_count; /* Reference count (atomic) */
gchar* name; /* Provider name (e.g., "openai", "perplexity") */
gchar* api_url; /* API endpoint URL */
gchar* org_id; /* Optional organization ID */
gchar* project_id; /* Optional project ID (for some providers) */
gchar* default_model; /* Default model for this provider */
GHashTable* settings; /* Extensible per-provider settings (e.g., tools=enabled, search=disabled) */
GList* models; /* Cached models (gchar*) */
gboolean models_fresh; /* Whether models cache is current */
gint32 ref_count; /* Reference count (atomic) */
} AIProvider;
/**
@@ -97,6 +100,15 @@ GList* ai_list_providers(void);
*/
gchar* ai_providers_find(const char* const search_str, gboolean previous, void* context);
/**
* Find a model name for autocomplete.
* @param search_str The search string
* @param previous Whether to go to previous match
* @param context ProfAiWin* for the current session
* @return Model name, or NULL if not found
*/
gchar* ai_models_find(const char* const search_str, gboolean previous, void* context);
/**
* Get the API key for a provider.
* @param provider_name The provider name
@@ -111,6 +123,51 @@ gchar* ai_get_provider_key(const gchar* provider_name);
*/
void ai_set_provider_key(const gchar* provider_name, const gchar* api_key);
/**
* Set the default model for a provider.
* @param provider_name The provider name
* @param model The default model name
*/
void ai_set_provider_default_model(const gchar* provider_name, const gchar* model);
/**
* Get the default model for a provider.
* @param provider_name The provider name
* @return The default model name (caller must not free), or NULL if not set
*/
const gchar* ai_get_provider_default_model(const gchar* provider_name);
/**
* Set a custom setting for a provider.
* @param provider_name The provider name
* @param setting The setting key (e.g., "tools", "search")
* @param value The setting value
*/
void ai_set_provider_setting(const gchar* provider_name, const gchar* setting, const gchar* value);
/**
* Get a custom setting for a provider.
* @param provider_name The provider name
* @param setting The setting key
* @return The setting value (caller must free), or NULL if not set
*/
gchar* ai_get_provider_setting(const gchar* provider_name, const gchar* setting);
/**
* Fetch available models from a provider's API.
* @param provider_name The provider name
* @param user_data User data (ProfAiWin* for UI display)
* @return TRUE if the request was successfully queued, FALSE otherwise
*/
gboolean ai_fetch_models(const gchar* provider_name, gpointer user_data);
/**
* Check if models cache is fresh for a provider.
* @param provider_name The provider name
* @return TRUE if models are fresh, FALSE otherwise
*/
gboolean ai_models_are_fresh(const gchar* provider_name);
/* ========================================================================
* Session Management
* ======================================================================== */