using System;
using System.Collections.Generic;
using System.IO;
using System.Linq;
using System.Net.Http;
using System.Text;
using System.Threading.Tasks;
using Newtonsoft.Json.Linq;
using SwarmUI.Accounts;
using SwarmUI.Utils;
namespace Mrleo1nid.SwarmAssistent;
/// Training samples, dataset import/export, Ollama Modelfile builder.
public partial class SwarmAssistentExtension
{
static string TrainingRoot()
{
string root = Path.Combine(DataRoot(), "Assistent", "training");
Directory.CreateDirectory(root);
Directory.CreateDirectory(Path.Combine(root, "datasets"));
Directory.CreateDirectory(Path.Combine(root, "jobs"));
Directory.CreateDirectory(Path.Combine(root, "adapters"));
return root;
}
public async Task AssistentListTrainSamples(Session session, string status = null, string persona = null, int limit = 200)
{
await Task.CompletedTask;
try
{
List list = Memory.ListTrainSamples(status, persona, null, limit);
return new JObject
{
["success"] = true,
["samples"] = new JArray(list),
["approved"] = Memory.CountTrainSamples("approved"),
["draft"] = Memory.CountTrainSamples("draft"),
["total"] = Memory.CountTrainSamples(null),
};
}
catch (Exception ex)
{
return new JObject { ["error"] = ex.Message };
}
}
public async Task AssistentUpsertTrainSample(Session session, JObject raw)
{
await Task.CompletedTask;
if (raw is null)
{
return new JObject { ["error"] = "body required" };
}
try
{
JObject saved = Memory.UpsertTrainSample(raw);
JObject full = Memory.GetTrainSample(saved["id"]?.ToString()) ?? raw;
await TryAutoLinkTrainSample(session, full, raw);
return new JObject { ["success"] = true, ["sample"] = full };
}
catch (Exception ex)
{
return new JObject { ["error"] = ex.Message };
}
}
public async Task AssistentDeleteTrainSample(Session session, string id)
{
await Task.CompletedTask;
if (string.IsNullOrWhiteSpace(id))
{
return new JObject { ["error"] = "id required" };
}
JObject sample = Memory.GetTrainSample(id.Trim());
bool ok = Memory.DeleteTrainSample(id.Trim());
if (ok && sample is not null)
{
Memory.UnlinkTrainSampleFromAgent(sample);
}
return new JObject { ["success"] = true, ["deleted"] = ok };
}
public async Task AssistentBuildDatasetFromChats(Session session, bool approved_only = false)
{
await Task.CompletedTask;
try
{
List chats = Memory.ListChats(withMessages: true, limit: AssistentMemory.MaxChatsStored);
int added = 0;
foreach (JObject chat in chats)
{
JArray messages = chat["messages"] as JArray ?? [];
for (int i = 0; i < messages.Count - 1; i++)
{
if (messages[i] is not JObject u || messages[i + 1] is not JObject a)
{
continue;
}
if (!string.Equals(u["role"]?.ToString(), "user", StringComparison.OrdinalIgnoreCase))
{
continue;
}
if (!string.Equals(a["role"]?.ToString(), "assistant", StringComparison.OrdinalIgnoreCase))
{
continue;
}
Memory.UpsertTrainSample(new JObject
{
["source"] = "chat",
["chat_id"] = chat["id"],
["persona"] = a["persona"] ?? u["persona"],
["pack"] = a["pack"] ?? u["pack"],
["status"] = approved_only ? "approved" : "draft",
["messages"] = new JArray { u.DeepClone(), a.DeepClone() },
});
added++;
}
}
return new JObject { ["success"] = true, ["added"] = added };
}
catch (Exception ex)
{
return new JObject { ["error"] = ex.Message };
}
}
public async Task AssistentImportDataset(Session session, JObject raw)
{
await Task.CompletedTask;
string content = raw?["content"]?.ToString();
string format = raw?["format"]?.ToString() ?? "auto";
if (string.IsNullOrWhiteSpace(content))
{
return new JObject { ["error"] = "content required" };
}
try
{
int imported = 0;
string fmt = (format ?? "auto").Trim().ToLowerInvariant();
List records = ParseDatasetContent(content, fmt);
foreach (JObject rec in records)
{
JArray messages = rec["messages"] as JArray;
if (messages is null || messages.Count == 0)
{
continue;
}
Memory.UpsertTrainSample(new JObject
{
["source"] = "import",
["messages"] = messages,
["status"] = "draft",
});
imported++;
}
return new JObject { ["success"] = true, ["imported"] = imported };
}
catch (Exception ex)
{
return new JObject { ["error"] = ex.Message };
}
}
static List ParseDatasetContent(string content, string format)
{
List list = [];
string trimmed = content.Trim();
if (trimmed.StartsWith('['))
{
JArray arr = JArray.Parse(trimmed);
foreach (JToken t in arr)
{
if (t is JObject o)
{
JArray msgs = ExtractMessagesFromRecord(o);
if (msgs != null)
{
list.Add(new JObject { ["messages"] = msgs });
}
}
}
return list;
}
if (format == "csv" || LooksLikeCsv(trimmed))
{
return ParseCsvDataset(trimmed);
}
foreach (string line in trimmed.Split('\n'))
{
string ln = line.Trim();
if (string.IsNullOrWhiteSpace(ln))
{
continue;
}
try
{
JObject o = JObject.Parse(ln);
JArray msgs = ExtractMessagesFromRecord(o);
if (msgs != null)
{
list.Add(new JObject { ["messages"] = msgs });
}
}
catch
{
// skip bad line
}
}
return list;
}
static bool LooksLikeCsv(string s) => s.Contains(',') && s.Contains('\n') && !s.TrimStart().StartsWith('{');
static List ParseCsvDataset(string csv)
{
List list = [];
string[] lines = csv.Split('\n').Select(l => l.Trim()).Where(l => l.Length > 0).ToArray();
if (lines.Length < 2)
{
return list;
}
string[] headers = lines[0].Split(',').Select(h => h.Trim().Trim('"')).ToArray();
int promptIdx = Array.FindIndex(headers, h => h.Equals("prompt", StringComparison.OrdinalIgnoreCase) || h.Equals("question", StringComparison.OrdinalIgnoreCase) || h.Equals("instruction", StringComparison.OrdinalIgnoreCase));
int respIdx = Array.FindIndex(headers, h => h.Equals("response", StringComparison.OrdinalIgnoreCase) || h.Equals("answer", StringComparison.OrdinalIgnoreCase) || h.Equals("output", StringComparison.OrdinalIgnoreCase) || h.Equals("completion", StringComparison.OrdinalIgnoreCase));
if (promptIdx < 0 || respIdx < 0)
{
return list;
}
for (int i = 1; i < lines.Length; i++)
{
string[] cols = SplitCsvLine(lines[i]);
if (cols.Length <= Math.Max(promptIdx, respIdx))
{
continue;
}
list.Add(new JObject
{
["messages"] = new JArray
{
new JObject { ["role"] = "user", ["content"] = cols[promptIdx] },
new JObject { ["role"] = "assistant", ["content"] = cols[respIdx] },
},
});
}
return list;
}
static string[] SplitCsvLine(string line)
{
List parts = [];
StringBuilder cur = new();
bool inQ = false;
foreach (char c in line)
{
if (c == '"')
{
inQ = !inQ;
continue;
}
if (c == ',' && !inQ)
{
parts.Add(cur.ToString().Trim());
cur.Clear();
continue;
}
cur.Append(c);
}
parts.Add(cur.ToString().Trim());
return parts.ToArray();
}
static JArray ExtractMessagesFromRecord(JObject o)
{
if (o["messages"] is JArray msgs)
{
return NormalizeMessagesArray(msgs);
}
if (o["conversations"] is JArray conv)
{
return ConvertHfRowToMessages(new JObject { ["conversations"] = conv }, new JObject { ["kind"] = "conversations" }, null);
}
if (o["instruction"] != null && o["output"] != null)
{
return ConvertHfRowToMessages(o, new JObject { ["kind"] = "alpaca" }, null);
}
if (o["prompt"] != null && (o["response"] != null || o["completion"] != null))
{
return ConvertHfRowToMessages(o, new JObject { ["kind"] = "prompt_response", ["response_col"] = o["response"] != null ? "response" : "completion" }, null);
}
return null;
}
public async Task AssistentExportDataset(Session session, string status = "approved", string format = "jsonl")
{
await Task.CompletedTask;
try
{
List samples = Memory.ListTrainSamples(status, null, null, 5000);
StringBuilder sb = new();
int exported = 0;
foreach (JObject s in samples)
{
JArray messages = s["messages"] as JArray ?? [];
if (messages.Count == 0)
{
continue;
}
exported++;
if (string.Equals(format, "sharegpt", StringComparison.OrdinalIgnoreCase))
{
JArray conv = [];
foreach (JToken m in messages)
{
if (m is not JObject mo)
{
continue;
}
string role = mo["role"]?.ToString() ?? "user";
conv.Add(new JObject
{
["from"] = role == "assistant" ? "gpt" : "human",
["value"] = mo["content"]?.ToString() ?? "",
});
}
sb.AppendLine(new JObject { ["conversations"] = conv }.ToString(Newtonsoft.Json.Formatting.None));
}
else
{
sb.AppendLine(new JObject { ["messages"] = messages }.ToString(Newtonsoft.Json.Formatting.None));
}
}
string path = Path.Combine(TrainingRoot(), "datasets", $"export_{DateTimeOffset.UtcNow.ToUnixTimeMilliseconds()}.jsonl");
await File.WriteAllTextAsync(path, sb.ToString(), Encoding.UTF8);
return new JObject { ["success"] = true, ["path"] = path, ["count"] = exported, ["content"] = sb.ToString() };
}
catch (Exception ex)
{
return new JObject { ["error"] = ex.Message };
}
}
public async Task AssistentCreateOllamaModel(Session session, JObject raw)
{
if (raw is null)
{
return new JObject { ["error"] = "body required" };
}
string baseUrl = NormalizeBaseUrl(raw["base_url"]?.ToString());
string baseModel = raw["base_model"]?.ToString()?.Trim();
string name = raw["name"]?.ToString()?.Trim();
string system = raw["system"]?.ToString() ?? "";
int shots = raw["shots"]?.Value() ?? 8;
if (string.IsNullOrWhiteSpace(baseModel) || string.IsNullOrWhiteSpace(name))
{
return new JObject { ["error"] = "base_model and name required" };
}
if (string.IsNullOrWhiteSpace(system))
{
string persona = AssistentConfig.SafeId(raw["persona"]?.ToString()) ?? Config.DefaultPersonaId();
system = Config.LoadCorePrompt(persona) + "\n\n" + Config.RenderIdentityBlock(persona, includeAllShelves: true);
}
StringBuilder mf = new();
mf.AppendLine($"FROM {baseModel}");
mf.AppendLine($"SYSTEM \"\"\"{system}\"\"\"");
List samples = Memory.ListTrainSamples("approved", null, null, Math.Clamp(shots, 0, 32));
foreach (JObject s in samples.Take(shots))
{
JArray messages = s["messages"] as JArray ?? [];
foreach (JToken m in messages)
{
if (m is not JObject mo)
{
continue;
}
string role = mo["role"]?.ToString() ?? "user";
string content = mo["content"]?.ToString() ?? "";
if (string.IsNullOrWhiteSpace(content))
{
continue;
}
mf.AppendLine($"MESSAGE {role} \"\"\"{content.Replace("\"\"\"", "\"\"\"\"\"\"\"")}\"\"\"");
}
}
if (raw["num_ctx"] != null)
{
mf.AppendLine($"PARAMETER num_ctx {raw["num_ctx"]}");
}
if (raw["temperature"] != null)
{
mf.AppendLine($"PARAMETER temperature {raw["temperature"]}");
}
string modelfilePath = Path.Combine(TrainingRoot(), "jobs", $"modelfile_{DateTimeOffset.UtcNow.ToUnixTimeMilliseconds()}.Modelfile");
Directory.CreateDirectory(Path.GetDirectoryName(modelfilePath)!);
await File.WriteAllTextAsync(modelfilePath, mf.ToString(), Encoding.UTF8);
try
{
JObject payload = new()
{
["name"] = name,
["modelfile"] = mf.ToString(),
["stream"] = false,
};
using StringContent content = new(payload.ToString(Newtonsoft.Json.Formatting.None), Encoding.UTF8, "application/json");
using HttpResponseMessage resp = await HttpClient.PostAsync($"{baseUrl}/api/create", content);
string body = await resp.Content.ReadAsStringAsync();
if (!resp.IsSuccessStatusCode)
{
return new JObject { ["error"] = $"Ollama create HTTP {(int)resp.StatusCode}: {Clip(body, 400)}", ["modelfile_path"] = modelfilePath };
}
return new JObject { ["success"] = true, ["name"] = name, ["modelfile_path"] = modelfilePath, ["ollama"] = body };
}
catch (Exception ex)
{
return new JObject { ["error"] = ex.Message, ["modelfile_path"] = modelfilePath };
}
}
public async Task AssistentGetTrainJob(Session session, string id = null)
{
await Task.CompletedTask;
JObject job = string.IsNullOrWhiteSpace(id) ? Memory.GetActiveTrainJob() : Memory.GetTrainJob(id);
if (job is not null && TrainingJobManager.IsRunning && string.Equals(job["id"]?.ToString(), TrainingJobManager.CurrentJobId, StringComparison.OrdinalIgnoreCase))
{
JObject live = TrainingJobManager.GetProgress();
job["progress_json"] = live.ToString(Newtonsoft.Json.Formatting.None);
job["status"] = live["status"]?.ToString() ?? job["status"];
}
return new JObject
{
["success"] = true,
["job"] = job,
["training_active"] = TrainingJobManager.IsRunning,
["progress"] = TrainingJobManager.IsRunning ? TrainingJobManager.GetProgress() : null,
};
}
public async Task AssistentSaveRunnerSettings(Session session, JObject settings)
{
await Task.CompletedTask;
if (settings is null)
{
return new JObject { ["error"] = "settings required" };
}
Config.SaveTrainingRunner(settings);
return new JObject { ["success"] = true };
}
public async Task AssistentGetRunnerSettings(Session session)
{
await Task.CompletedTask;
return new JObject { ["success"] = true, ["settings"] = Config.LoadTrainingRunner() };
}
}