using System;
using System.Net.Http;
using System.Text;
using System.Threading.Tasks;
using Newtonsoft.Json.Linq;
using SwarmUI.Accounts;
using SwarmUI.Utils;
namespace Mrleo1nid.SwarmAssistent;
/// VRAM handover between Ollama and the image backend: park the chat model before
/// Generate, warm it again once the user is back in the chat. Embed / memory models are never
/// parked — they are tiny and reloading them stalls every retrieve.
/// Warm is a no-op when Ollama already has the chat model resident (/api/ps) so post-Generate
/// reload is cheap when Krea did not evict VL.
public partial class SwarmAssistentExtension
{
const string WarmKeepAlive = "15m";
/// Unloads the chat model from VRAM (keep_alive: 0) so Krea 2 gets the whole GPU.
public async Task AssistentParkLlm(Session session, string baseUrl, string model)
{
string root = NormalizeBaseUrl(baseUrl);
string name = (model ?? "").Trim();
if (string.IsNullOrWhiteSpace(name))
{
return new JObject { ["error"] = "model is required" };
}
if (LooksLikeEmbedModel(name))
{
return new JObject { ["success"] = true, ["parked"] = false, ["skipped"] = "memory model — never parked" };
}
// Already unloaded — skip the keep_alive:0 round-trip.
if (!await IsOllamaModelResident(root, name))
{
return new JObject
{
["success"] = true,
["parked"] = false,
["skipped"] = "not_resident",
["model"] = name,
["base_url"] = root,
};
}
JObject generate = new()
{
["model"] = name,
["prompt"] = "",
["stream"] = false,
["keep_alive"] = 0,
};
(bool ok, string body) = await PostOllamaJson(root, "/api/generate", generate);
if (!ok)
{
// Older Ollama builds only unload through /api/chat.
JObject chat = new()
{
["model"] = name,
["messages"] = new JArray(),
["stream"] = false,
["keep_alive"] = 0,
};
(ok, body) = await PostOllamaJson(root, "/api/chat", chat);
}
if (!ok)
{
Logs.Debug($"AssistentParkLlm {name}: {Clip(body, 200)}");
return new JObject { ["success"] = true, ["parked"] = false, ["note"] = Clip(body, 200) };
}
return new JObject { ["success"] = true, ["parked"] = true, ["model"] = name, ["base_url"] = root };
}
/// Single-token chat so the model is resident again by the time the user types.
/// Skips the load when /api/ps already lists the model (avoids ~20–30s no-op warm).
public async Task AssistentWarmLlm(Session session, string baseUrl, string model)
{
string root = NormalizeBaseUrl(baseUrl);
string name = (model ?? "").Trim();
if (string.IsNullOrWhiteSpace(name))
{
return new JObject { ["error"] = "model is required" };
}
if (LooksLikeEmbedModel(name))
{
return new JObject { ["success"] = true, ["warmed"] = false, ["skipped"] = "memory model" };
}
if (await IsOllamaModelResident(root, name))
{
return new JObject
{
["success"] = true,
["warmed"] = true,
["skipped"] = "already_resident",
["model"] = name,
["keep_alive"] = WarmKeepAlive,
};
}
int numCtx = CfgInt("num_ctx", DefaultNumCtxFallback);
JObject payload = new()
{
["model"] = name,
["stream"] = false,
["messages"] = new JArray
{
new JObject { ["role"] = "user", ["content"] = "ok" },
},
["options"] = new JObject
{
["num_ctx"] = numCtx,
["num_predict"] = 1,
},
["keep_alive"] = WarmKeepAlive,
};
(bool ok, string body) = await PostOllamaJson(root, "/api/chat", payload);
if (!ok)
{
Logs.Debug($"AssistentWarmLlm {name}: {Clip(body, 200)}");
return new JObject { ["success"] = true, ["warmed"] = false, ["note"] = Clip(body, 200) };
}
return new JObject
{
["success"] = true,
["warmed"] = true,
["model"] = name,
["num_ctx"] = numCtx,
["keep_alive"] = WarmKeepAlive,
};
}
/// True when Ollama /api/ps lists (or a matching tag).
async Task IsOllamaModelResident(string root, string model)
{
if (string.IsNullOrWhiteSpace(root) || string.IsNullOrWhiteSpace(model))
{
return false;
}
try
{
using HttpResponseMessage resp = await HttpClient.GetAsync($"{root}/api/ps");
if (!resp.IsSuccessStatusCode)
{
return false;
}
string body = await resp.Content.ReadAsStringAsync();
JObject parsed = JObject.Parse(body);
JArray models = parsed["models"] as JArray ?? [];
foreach (JToken m in models)
{
string name = m["name"]?.ToString() ?? m["model"]?.ToString() ?? "";
if (string.IsNullOrWhiteSpace(name))
{
continue;
}
if (string.Equals(name, model, StringComparison.OrdinalIgnoreCase))
{
return true;
}
// Tags may differ by :latest vs bare name.
if (name.StartsWith(model + ":", StringComparison.OrdinalIgnoreCase)
|| model.StartsWith(name + ":", StringComparison.OrdinalIgnoreCase))
{
return true;
}
}
return false;
}
catch (Exception ex)
{
Logs.Debug($"IsOllamaModelResident: {ex.Message}");
return false;
}
}
static async Task<(bool ok, string body)> PostOllamaJson(string root, string route, JObject payload)
{
try
{
using StringContent content = new(payload.ToString(Newtonsoft.Json.Formatting.None), Encoding.UTF8, "application/json");
using HttpResponseMessage resp = await HttpClient.PostAsync($"{root}{route}", content);
string body = await resp.Content.ReadAsStringAsync();
return (resp.IsSuccessStatusCode, body);
}
catch (Exception ex)
{
return (false, ex.Message);
}
}
}