Files
profanity/src/ai/ai_client.c

840 lines
26 KiB
C

/*
* ai_client.c - AI client for interacting with OpenAI-compatible API providers
*
* Supports multiple providers (OpenAI, Perplexity, etc.) with per-provider
* API keys, custom endpoints, and model selection.
*
* Copyright (C) 2026 CProof Developers
*
* SPDX-License-Identifier: GPL-3.0-or-later WITH OpenSSL-exception
*/
// vim: expandtab:ts=4:sts=4:sw=4
#include "ai_client.h"
#include "common.h"
#include "log.h"
#include "config/preferences.h"
#include "profanity.h"
#include "tools/autocomplete.h"
#include "ui/ui.h"
#include "ui/window_list.h"
#include <curl/curl.h>
#include <glib.h>
#include <json-glib/json-glib.h>
#include <string.h>
/* Default providers */
#define DEFAULT_OPENAI_URL "https://api.openai.com/"
#define DEFAULT_PERPLEXITY_URL "https://api.perplexity.ai/"
/* Global state */
static GHashTable* providers = NULL;
static GHashTable* provider_keys = NULL;
static Autocomplete providers_ac = NULL;
/* ========================================================================
* Curl helpers
* ======================================================================== */
struct curl_response_t
{
gchar* data;
size_t size;
};
static size_t
_write_callback(void* ptr, size_t size, size_t nmemb, void* userdata)
{
size_t realsize = size * nmemb;
struct curl_response_t* res = (struct curl_response_t*)userdata;
/* Limit response size to 10MB to prevent OOM */
if (res->size + realsize > 10 * 1024 * 1024) {
log_error("[AI-THREAD] Response too large, truncating");
return realsize;
}
gchar* new_data = g_realloc(res->data, res->size + realsize + 1);
if (!new_data) {
log_error("[AI-THREAD] Failed to allocate memory for response");
return realsize;
}
res->data = new_data;
memcpy(res->data + res->size, ptr, realsize);
res->size += realsize;
res->data[res->size] = '\0';
return realsize;
}
/* ========================================================================
* Thread-safe UI display helpers
* ======================================================================== */
/**
* Validate aiwin pointer by checking if it still exists in the window list.
* Detects if the original window was freed (TOCTOU protection).
* Returns the validated pointer if valid, NULL if window was closed.
*/
static ProfAiWin*
_aiwin_validate(gpointer user_data)
{
if (!user_data) {
return NULL;
}
if (!wins_ai_exists((ProfAiWin*)user_data)) {
log_warning("[AI-THREAD] aiwin=%p no longer exists — window was closed", (void*)user_data);
return NULL;
}
/* Pointer is valid */
return (ProfAiWin*)user_data;
}
/**
* Display an error message in the AI window (thread-safe).
* @param user_data The original user_data pointer (may be NULL)
* @param error_msg The error message to display
*/
static void
_aiwin_display_error(gpointer user_data, const gchar* error_msg)
{
ProfAiWin* aiwin = _aiwin_validate(user_data);
if (!aiwin) {
log_warning("[AI-THREAD] Cannot display error: aiwin is invalid or window was closed (msg: %s)", error_msg);
return;
}
pthread_mutex_lock(&lock);
aiwin_display_error(aiwin, error_msg);
pthread_mutex_unlock(&lock);
}
/**
* Display a response message in the AI window (thread-safe).
* @param user_data The original user_data pointer (may be NULL)
* @param response The response message to display
*/
static void
_aiwin_display_response(gpointer user_data, const gchar* response)
{
ProfAiWin* aiwin = _aiwin_validate(user_data);
if (!aiwin) {
log_warning("[AI-THREAD] Cannot display response: aiwin is invalid or window was closed");
return;
}
pthread_mutex_lock(&lock);
aiwin_display_response(aiwin, response);
pthread_mutex_unlock(&lock);
}
/* ========================================================================
* JSON helpers
* ======================================================================== */
gchar*
ai_json_escape(const gchar* str)
{
if (!str)
return g_strdup("");
GString* result = g_string_new("");
for (const gchar* p = str; *p; p++) {
switch (*p) {
case '"':
g_string_append(result, "\\\"");
break;
case '\\':
g_string_append(result, "\\\\");
break;
case '\b':
g_string_append(result, "\\b");
break;
case '\f':
g_string_append(result, "\\f");
break;
case '\n':
g_string_append(result, "\\n");
break;
case '\r':
g_string_append(result, "\\r");
break;
case '\t':
g_string_append(result, "\\t");
break;
default:
g_string_append_c(result, *p);
break;
}
}
return g_string_free(result, FALSE);
}
/* ========================================================================
* Provider Management
* ======================================================================== */
static AIProvider*
ai_provider_new(const gchar* name, const gchar* api_url, const gchar* org_id)
{
AIProvider* provider = g_new0(AIProvider, 1);
provider->name = g_strdup(name);
provider->api_url = g_strdup(api_url ? api_url : "");
provider->org_id = g_strdup(org_id);
provider->project_id = NULL;
provider->models = NULL;
provider->ref_count = 1;
return provider;
}
static AIProvider*
ai_provider_ref(AIProvider* provider)
{
if (provider) {
g_atomic_int_inc(&provider->ref_count);
}
return provider;
}
void
ai_provider_unref(AIProvider* provider)
{
if (!provider)
return;
if (!g_atomic_int_dec_and_test(&provider->ref_count)) {
return;
}
g_free(provider->name);
g_free(provider->api_url);
g_free(provider->org_id);
g_free(provider->project_id);
GList* curr = provider->models;
while (curr) {
g_free(curr->data);
curr = g_list_next(curr);
}
g_list_free(provider->models);
g_free(provider);
}
static void
ai_load_keys(void)
{
if (!provider_keys) {
provider_keys = g_hash_table_new_full(g_str_hash, g_str_equal, g_free, g_free);
}
GList* tokens = prefs_ai_list_tokens();
if (!tokens) {
return;
}
GList* curr = tokens;
while (curr) {
gchar* provider = (gchar*)curr->data;
gchar* key = prefs_ai_get_token(provider);
if (key && strlen(key) > 0) {
g_hash_table_insert(provider_keys, g_strdup(provider), g_strdup(key));
}
g_free(key);
curr = g_list_next(curr);
}
prefs_free_ai_tokens(tokens);
log_info("Loaded %d saved API keys from config", g_hash_table_size(provider_keys));
}
void
ai_client_init(void)
{
if (providers)
return; /* Already initialized */
curl_global_init(CURL_GLOBAL_ALL);
/* Create hash tables */
providers = g_hash_table_new_full(g_str_hash, g_str_equal, g_free, (GDestroyNotify)ai_provider_unref);
provider_keys = g_hash_table_new_full(g_str_hash, g_str_equal, g_free, g_free);
/* Create autocomplete for provider names */
providers_ac = autocomplete_new();
/* Add default providers */
ai_add_provider("openai", DEFAULT_OPENAI_URL, NULL);
ai_add_provider("perplexity", DEFAULT_PERPLEXITY_URL, NULL);
/* Load saved API keys from config */
ai_load_keys();
log_info("AI client initialized with default providers: openai, perplexity");
}
void
ai_client_shutdown(void)
{
if (!providers)
return;
g_hash_table_destroy(providers);
g_hash_table_destroy(provider_keys);
providers = NULL;
provider_keys = NULL;
if (providers_ac) {
autocomplete_free(providers_ac);
providers_ac = NULL;
}
curl_global_cleanup();
log_info("AI client shutdown");
}
AIProvider*
ai_get_provider(const gchar* name)
{
if (!name || !providers)
return NULL;
return g_hash_table_lookup(providers, name);
}
AIProvider*
ai_add_provider(const gchar* name, const gchar* api_url, const gchar* org_id)
{
if (!name || !api_url)
return NULL;
if (!providers) {
ai_client_init();
}
/* Check if provider already exists */
AIProvider* existing = g_hash_table_lookup(providers, name);
if (existing) {
/* Update existing provider */
g_free(existing->api_url);
existing->api_url = g_strdup(api_url);
g_free(existing->org_id);
existing->org_id = g_strdup(org_id);
log_info("Updated provider: %s", name);
return ai_provider_ref(existing);
}
/* Create new provider (ref_count=1 owned by hash table) */
AIProvider* provider = ai_provider_new(name, api_url, org_id);
g_hash_table_insert(providers, g_strdup(name), provider);
/* Sync autocomplete */
autocomplete_add(providers_ac, name);
log_info("Added provider: %s (URL: %s)", name, api_url);
return provider; /* Caller gets non-owning pointer; hash table owns ref */
}
gboolean
ai_remove_provider(const gchar* name)
{
if (!name || !providers)
return FALSE;
/* Sync autocomplete before removing */
autocomplete_remove(providers_ac, name);
return g_hash_table_remove(providers, name);
}
GList*
ai_list_providers(void)
{
if (!providers)
return NULL;
GList* result = NULL;
GHashTableIter iter;
gpointer key, value;
g_hash_table_iter_init(&iter, providers);
while (g_hash_table_iter_next(&iter, &key, &value)) {
result = g_list_append(result, value);
}
return result;
}
/* ========================================================================
* Provider autocomplete state
* ======================================================================== */
/* Stateful provider name finder with case-sensitive matching.
* Maintains position in sorted list for deterministic tab-completion cycling.
* The autocomplete is kept in sync via ai_add_provider/ai_remove_provider. */
gchar*
ai_providers_find(const char* const search_str, gboolean previous, void* context)
{
/* Initialize autocomplete on first use */
if (!providers_ac) {
providers_ac = autocomplete_new();
}
/* NULL search_str is treated as empty string for cycling */
const char* effective_search = (search_str != NULL) ? search_str : "";
/* Use stateful autocomplete */
return autocomplete_complete(providers_ac, effective_search, FALSE, previous);
}
gchar*
ai_get_provider_key(const gchar* provider_name)
{
if (!provider_name || !provider_keys)
return NULL;
gchar* key = g_hash_table_lookup(provider_keys, provider_name);
return g_strdup(key);
}
void
ai_set_provider_key(const gchar* provider_name, const gchar* api_key)
{
if (!provider_name)
return;
if (!provider_keys) {
provider_keys = g_hash_table_new_full(g_str_hash, g_str_equal, g_free, g_free);
}
if (api_key) {
g_hash_table_insert(provider_keys, g_strdup(provider_name), g_strdup(api_key));
/* Persist to config file */
prefs_ai_set_token(provider_name, api_key);
log_info("API key set for provider: %s", provider_name);
} else {
g_hash_table_remove(provider_keys, provider_name);
/* Remove from config file */
prefs_ai_remove_token(provider_name);
log_info("API key removed for provider: %s", provider_name);
}
}
/* ========================================================================
* Session Management
* ======================================================================== */
AISession*
ai_session_create(const gchar* provider_name, const gchar* model)
{
if (!provider_name || !model)
return NULL;
AIProvider* provider = ai_get_provider(provider_name);
if (!provider) {
log_error("Provider not found: %s", provider_name);
return NULL;
}
AISession* session = g_new0(AISession, 1);
session->provider_name = g_strdup(provider_name);
session->provider = ai_provider_ref(provider);
session->model = g_strdup(model);
session->api_key = ai_get_provider_key(provider_name);
session->history = NULL;
session->ref_count = 1;
log_info("AI session created: %s/%s", provider_name, model);
return session;
}
AISession*
ai_session_ref(AISession* session)
{
if (session) {
g_atomic_int_inc(&session->ref_count);
}
return session;
}
void
ai_session_unref(AISession* session)
{
if (!session)
return;
if (!g_atomic_int_dec_and_test(&session->ref_count)) {
return;
}
g_free(session->provider_name);
ai_provider_unref(session->provider);
g_free(session->model);
g_free(session->api_key);
GList* curr = session->history;
while (curr) {
AIMessage* msg = curr->data;
g_free(msg->role);
g_free(msg->content);
g_free(msg);
curr = g_list_next(curr);
}
g_list_free(session->history);
g_free(session);
log_debug("AI session destroyed");
}
void
ai_session_add_message(AISession* session, const gchar* role, const gchar* content)
{
if (!session || !role || !content)
return;
AIMessage* msg = g_new0(AIMessage, 1);
msg->role = g_strdup(role);
msg->content = g_strdup(content);
session->history = g_list_append(session->history, msg);
log_debug("Added %s message to session (total: %d)", role, g_list_length(session->history));
}
void
ai_session_clear_history(AISession* session)
{
if (!session)
return;
GList* curr = session->history;
while (curr) {
AIMessage* msg = curr->data;
g_free(msg->role);
g_free(msg->content);
g_free(msg);
curr = g_list_next(curr);
}
g_list_free(session->history);
session->history = NULL;
log_info("AI session history cleared");
}
const gchar*
ai_session_get_model(AISession* session)
{
if (!session)
return NULL;
return session->model;
}
void
ai_session_set_model(AISession* session, const gchar* model)
{
if (!session || !model)
return;
g_free(session->model);
session->model = g_strdup(model);
log_info("Session model changed to: %s", model);
}
/* ========================================================================
* API Request Handling
* ======================================================================== */
static gchar*
_build_json_payload(AISession* session, const gchar* prompt)
{
/* OpenAI-compatible Responses API format:
* {"model": "...", "input": [...], "stream": false, "store": false}
* store:false prevents providers from storing/using requests for training */
GString* messages_json = g_string_new("");
GList* curr = session->history;
while (curr) {
AIMessage* msg = curr->data;
auto_gchar gchar* escaped_content = ai_json_escape(msg->content);
auto_gchar gchar* escaped_role = ai_json_escape(msg->role);
if (messages_json->len > 0) {
g_string_append_c(messages_json, ',');
}
g_string_append_printf(messages_json, "{\"role\":\"%s\",\"content\":\"%s\"}",
escaped_role, escaped_content);
curr = g_list_next(curr);
}
/* Add the new user message */
auto_gchar gchar* escaped_prompt = ai_json_escape(prompt);
if (messages_json->len > 0) {
g_string_append_c(messages_json, ',');
}
g_string_append_printf(messages_json, "{\"role\":\"user\",\"content\":\"%s\"}", escaped_prompt);
auto_gchar gchar* escaped_model = ai_json_escape(session->model);
gchar* json_payload = g_strdup_printf(
"{\"model\":\"%s\",\"input\":[%s],\"stream\":false,\"store\":false}",
escaped_model, messages_json->str);
g_string_free(messages_json, TRUE);
return json_payload;
}
/** @brief Parse AI response JSON and extract content string.
*
* Tries multiple JSON paths to handle different provider response formats:
* 1. Perplexity /v1/agent: {"output":[{"content":[{"text":"..."}]}]}
* 2. OpenAI /v1/chat/completions: {"choices":[{"message":{"content":"..."}}]}
* 3. Fallback: root-level "text" field
* 4. Fallback: root-level "content" field
*
* @param response_json The JSON response string
* @return Newly allocated content string, or NULL on failure
*/
gchar*
ai_parse_response(const gchar* response_json)
{
if (!response_json || strlen(response_json) == 0) {
log_warning("[AI-THREAD] Empty or NULL response JSON");
return NULL;
}
/* Parse JSON using json-glib */
auto_gchar gchar* error_msg = NULL;
JsonParser* parser = json_parser_new();
if (!json_parser_load_from_data(parser, response_json, -1, &error_msg)) {
log_warning("[AI-THREAD] Failed to parse AI response JSON: %s", error_msg ? error_msg : "unknown error");
g_object_unref(parser);
return NULL;
}
JsonNode* root = json_parser_get_root(parser);
if (!root || !JSON_IS_OBJECT(json_node_get_object(root))) {
log_warning("[AI-THREAD] Invalid AI response JSON: root is not an object");
g_object_unref(parser);
return NULL;
}
JsonObject* root_obj = json_node_get_object(root);
gchar* content = NULL;
/* Try 1: Perplexity /v1/agent format
* {"output":[{"content":[{"text":"...","type":"output_text"}]}]} */
JsonNode* output_node = json_object_get_member(root_obj, "output");
if (output_node && JSON_IS_ARRAY(json_node_get_array(output_node))) {
JsonArray* output_arr = json_node_get_array(output_node);
if (json_array_get_length(output_arr) > 0) {
JsonObject* first_output = json_array_get_object_element(output_arr, 0);
JsonNode* content_node = json_object_get_member(first_output, "content");
if (content_node && JSON_IS_ARRAY(json_node_get_array(content_node))) {
JsonArray* content_arr = json_node_get_array(content_node);
if (json_array_get_length(content_arr) > 0) {
JsonObject* first_content = json_array_get_object_element(content_arr, 0);
JsonNode* text_node = json_object_get_member(first_content, "text");
if (text_node && JSON_IS_VALUE(text_node)) {
content = g_strdup(json_node_get_string(text_node));
}
}
}
}
}
/* Try 2: OpenAI /v1/chat/completions format
* {"choices":[{"message":{"content":"..."}}]} */
if (!content) {
JsonNode* choices_node = json_object_get_member(root_obj, "choices");
if (choices_node && JSON_IS_ARRAY(json_node_get_array(choices_node))) {
JsonArray* choices_arr = json_node_get_array(choices_node);
if (json_array_get_length(choices_arr) > 0) {
JsonObject* first_choice = json_array_get_object_element(choices_arr, 0);
JsonNode* message_node = json_object_get_member(first_choice, "message");
if (message_node && JSON_IS_OBJECT(json_node_get_object(message_node))) {
JsonObject* message = json_node_get_object(message_node);
JsonNode* content_member = json_object_get_member(message, "content");
if (content_member && JSON_IS_VALUE(content_member)) {
content = g_strdup(json_node_get_string(content_member));
}
}
}
}
}
/* Try 3: Fallback - look for "text" field anywhere in the response
* This handles edge cases where the response shape differs */
if (!content) {
JsonNode* text_node = json_object_get_member(root_obj, "text");
if (text_node && JSON_IS_VALUE(text_node)) {
content = g_strdup(json_node_get_string(text_node));
}
}
/* Try 4: Fallback - look for "content" field at root level */
if (!content) {
JsonNode* content_node = json_object_get_member(root_obj, "content");
if (content_node && JSON_IS_VALUE(content_node)) {
content = g_strdup(json_node_get_string(content_node));
}
}
g_object_unref(parser);
if (!content) {
log_warning("[AI-THREAD] Could not extract content from AI response");
}
return content;
}
static gpointer
_ai_request_thread(gpointer data)
{
log_debug("[AI-THREAD] Starting AI request thread");
/* Data is an array: [0]=session, [1]=prompt, [2]=user_data */
gpointer* args = (gpointer*)data;
AISession* session = (AISession*)args[0];
auto_gchar gchar* prompt = args[1];
gpointer user_data = args[2];
log_debug("[AI-THREAD] Session: %s/%s", session->provider_name, session->model);
log_debug("[AI-THREAD] API key length: %zu", session->api_key ? strlen(session->api_key) : 0);
/* Check for API key first */
if (!session->api_key || strlen(session->api_key) == 0) {
auto_gchar gchar* error_msg = g_strdup_printf("No API key set for provider '%s'. Use '/ai set token %s <key>' to configure.",
session->provider_name, session->provider_name);
log_error("AI request failed for %s/%s: %s", session->provider_name, session->model, error_msg);
_aiwin_display_error(user_data, error_msg);
g_free(args);
return NULL;
}
CURL* curl = curl_easy_init();
log_debug("[AI-THREAD] Curl initialized: %s", curl ? "OK" : "FAILED");
if (!curl) {
log_error("AI request failed for %s/%s: Failed to initialize curl",
session->provider_name, session->model);
_aiwin_display_error(user_data, "Failed to initialize curl.");
g_free(args);
return NULL;
}
/* Add user message to history FIRST */
ai_session_add_message(session, "user", prompt);
log_debug("[AI-THREAD] Added user message to history");
/* Build JSON payload (includes the message we just added) */
log_debug("[AI-THREAD] Building JSON payload...");
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);
}
/* Response buffer */
struct curl_response_t response;
response.data = g_new0(gchar, 1);
response.size = 0;
/* Configure request */
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);
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);
log_debug("[AI-THREAD] curl_easy_perform completed, res: %d", res);
/* Get HTTP response code */
long http_code = 0;
curl_easy_getinfo(curl, CURLINFO_RESPONSE_CODE, &http_code);
log_debug("[AI-THREAD] HTTP response code: %ld", http_code);
if (res != CURLE_OK) {
auto_gchar gchar* error_msg = g_strdup(curl_easy_strerror(res));
log_error("AI request failed for %s/%s: %s", session->provider_name, session->model, error_msg);
_aiwin_display_error(user_data, error_msg);
} else if (http_code >= 400) {
/* Handle HTTP errors */
log_debug("[AI-THREAD] HTTP error response body (%zu bytes): %s",
response.size, response.data ? response.data : "NULL");
auto_gchar gchar* error_msg = g_strdup_printf("HTTP %ld: %s", http_code,
response.data ? response.data : "Unknown error");
log_error("AI request failed for %s/%s: %s", session->provider_name, session->model, error_msg);
_aiwin_display_error(user_data, error_msg);
} else {
/* Parse response - transfer ownership to auto_gchar for cleanup */
log_debug("[AI-THREAD] Raw API response (%zu bytes): %s", response.size, response.data ? response.data : "NULL");
auto_gchar gchar* response_data = response.data;
response.data = NULL;
auto_gchar gchar* content = _parse_ai_response(response_data);
if (content) {
/* Add assistant response to history */
ai_session_add_message(session, "assistant", content);
_aiwin_display_response(user_data, content);
} else {
log_error("AI response parse failed for %s/%s: %.200s...",
session->provider_name, session->model, response_data);
_aiwin_display_error(user_data, "Failed to parse AI response.");
}
}
curl_slist_free_all(headers);
curl_easy_cleanup(curl);
g_free(args);
return NULL;
}
gboolean
ai_send_prompt(AISession* session, const gchar* prompt, gpointer user_data)
{
log_debug("[AI-PROMPT] ENTER: session=%p, prompt='%s', user_data=%p",
(void*)session, prompt, user_data);
if (!session || !prompt) {
log_error("[AI-PROMPT] FAIL: invalid session or prompt");
return FALSE;
}
/* Prepare thread arguments: [0]=session, [1]=prompt, [2]=user_data */
gpointer* args = g_new0(gpointer, 3);
args[0] = ai_session_ref(session);
args[1] = g_strdup(prompt);
args[2] = user_data;
log_debug("[AI-PROMPT] Prepared args, creating thread...");
GThread* thread = g_thread_new("ai_request", _ai_request_thread, args);
if (!thread) {
g_free(args[1]);
ai_session_unref(session);
g_free(args);
log_error("[AI-PROMPT] FAIL: g_thread_new returned NULL");
return FALSE;
}
log_debug("[AI-PROMPT] Thread created successfully: %p", (void*)thread);
g_thread_unref(thread);
return TRUE;
}