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() }; } }