using System; using System.Collections.Generic; using System.Linq; using Microsoft.Data.Sqlite; using Newtonsoft.Json.Linq; using SwarmUI.Utils; namespace Mrleo1nid.SwarmAssistent; /// Training datasets, samples, and job metadata in assistent.sqlite. public sealed partial class AssistentMemory { public const string KvHfDatasetCache = "hf_dataset_cache"; void EnsureTrainingSchema() { Exec( """ CREATE TABLE IF NOT EXISTS train_datasets ( id TEXT PRIMARY KEY, title TEXT NOT NULL DEFAULT '', created_at INTEGER NOT NULL, updated_at INTEGER NOT NULL, meta_json TEXT ); CREATE TABLE IF NOT EXISTS train_samples ( id TEXT PRIMARY KEY, dataset_id TEXT NOT NULL DEFAULT 'default', source TEXT NOT NULL DEFAULT 'manual', chat_id TEXT, persona TEXT, pack TEXT, hf_repo TEXT, messages_json TEXT NOT NULL DEFAULT '[]', status TEXT NOT NULL DEFAULT 'draft', created_at INTEGER NOT NULL, updated_at INTEGER NOT NULL, FOREIGN KEY(dataset_id) REFERENCES train_datasets(id) ON DELETE CASCADE ); CREATE INDEX IF NOT EXISTS idx_train_samples_status ON train_samples(status); CREATE INDEX IF NOT EXISTS idx_train_samples_dataset ON train_samples(dataset_id); CREATE TABLE IF NOT EXISTS train_jobs ( id TEXT PRIMARY KEY, kind TEXT NOT NULL, status TEXT NOT NULL DEFAULT 'pending', config_json TEXT, base_model TEXT, output_name TEXT, log_path TEXT, progress_json TEXT, created_at INTEGER NOT NULL, updated_at INTEGER NOT NULL, finished_at INTEGER ); CREATE INDEX IF NOT EXISTS idx_train_jobs_status ON train_jobs(status); """); if (!HasColumn("train_samples", "agent_linked")) { Exec("ALTER TABLE train_samples ADD COLUMN agent_linked INTEGER NOT NULL DEFAULT 0"); } using SqliteCommand cmd = _conn.CreateCommand(); cmd.CommandText = "INSERT OR IGNORE INTO train_datasets(id, title, created_at, updated_at) VALUES('default', 'Default', $u, $u)"; cmd.Parameters.AddWithValue("$u", DateTimeOffset.UtcNow.ToUnixTimeMilliseconds()); cmd.ExecuteNonQuery(); } public List ListTrainSamples(string status = null, string persona = null, string datasetId = null, int limit = 200) { lock (_lock) { EnsureOpen(); int take = Math.Clamp(limit, 1, 2000); List where = []; if (!string.IsNullOrWhiteSpace(status) && !string.Equals(status, "all", StringComparison.OrdinalIgnoreCase)) { where.Add("status = $status"); } if (!string.IsNullOrWhiteSpace(persona) && !string.Equals(persona, "all", StringComparison.OrdinalIgnoreCase)) { where.Add("persona = $persona"); } if (!string.IsNullOrWhiteSpace(datasetId)) { where.Add("dataset_id = $ds"); } bool hasAgentLinked = HasColumn("train_samples", "agent_linked"); string sql = hasAgentLinked ? "SELECT id, dataset_id, source, chat_id, persona, pack, hf_repo, messages_json, status, created_at, updated_at, agent_linked FROM train_samples" : "SELECT id, dataset_id, source, chat_id, persona, pack, hf_repo, messages_json, status, created_at, updated_at FROM train_samples"; if (where.Count > 0) { sql += " WHERE " + string.Join(" AND ", where); } sql += " ORDER BY updated_at DESC LIMIT $lim"; using SqliteCommand cmd = _conn.CreateCommand(); cmd.CommandText = sql; if (where.Any(w => w.Contains("$status"))) { cmd.Parameters.AddWithValue("$status", status.Trim()); } if (where.Any(w => w.Contains("$persona"))) { cmd.Parameters.AddWithValue("$persona", persona.Trim()); } if (where.Any(w => w.Contains("$ds"))) { cmd.Parameters.AddWithValue("$ds", datasetId.Trim()); } cmd.Parameters.AddWithValue("$lim", take); List list = []; using SqliteDataReader r = cmd.ExecuteReader(); while (r.Read()) { list.Add(ReadTrainSampleRow(r)); } return list; } } static JObject ReadTrainSampleRow(SqliteDataReader r) { JArray messages = []; try { messages = JArray.Parse(r.GetString(7)); } catch { // ignore } bool hasAgentLinked = r.FieldCount > 11; return new JObject { ["id"] = r.GetString(0), ["dataset_id"] = r.GetString(1), ["source"] = r.GetString(2), ["chat_id"] = r.IsDBNull(3) ? null : r.GetString(3), ["persona"] = r.IsDBNull(4) ? null : r.GetString(4), ["pack"] = r.IsDBNull(5) ? null : r.GetString(5), ["hf_repo"] = r.IsDBNull(6) ? null : r.GetString(6), ["messages"] = messages, ["status"] = r.GetString(8), ["createdAt"] = r.GetInt64(9), ["updatedAt"] = r.GetInt64(10), ["agent_linked"] = hasAgentLinked && !r.IsDBNull(11) && r.GetInt64(11) != 0, }; } public JObject GetTrainSample(string id) { if (string.IsNullOrWhiteSpace(id)) { return null; } lock (_lock) { EnsureOpen(); bool hasAgentLinked = HasColumn("train_samples", "agent_linked"); string sql = hasAgentLinked ? "SELECT id, dataset_id, source, chat_id, persona, pack, hf_repo, messages_json, status, created_at, updated_at, agent_linked FROM train_samples WHERE id = $id LIMIT 1" : "SELECT id, dataset_id, source, chat_id, persona, pack, hf_repo, messages_json, status, created_at, updated_at FROM train_samples WHERE id = $id LIMIT 1"; using SqliteCommand cmd = _conn.CreateCommand(); cmd.CommandText = sql; cmd.Parameters.AddWithValue("$id", id.Trim()); using SqliteDataReader r = cmd.ExecuteReader(); return r.Read() ? ReadTrainSampleRow(r) : null; } } public JObject UpsertTrainSample(JObject sample) { lock (_lock) { EnsureOpen(); long now = DateTimeOffset.UtcNow.ToUnixTimeMilliseconds(); string id = sample["id"]?.ToString()?.Trim(); if (string.IsNullOrWhiteSpace(id)) { id = $"ts_{now}_{Guid.NewGuid():N}"[..24]; } string datasetId = sample["dataset_id"]?.ToString()?.Trim() ?? "default"; JArray messages = sample["messages"] as JArray ?? []; using SqliteCommand cmd = _conn.CreateCommand(); cmd.CommandText = """ INSERT INTO train_samples(id, dataset_id, source, chat_id, persona, pack, hf_repo, messages_json, status, created_at, updated_at) VALUES($id, $ds, $src, $chat, $persona, $pack, $hf, $msg, $status, $c, $u) ON CONFLICT(id) DO UPDATE SET dataset_id = excluded.dataset_id, source = excluded.source, chat_id = excluded.chat_id, persona = excluded.persona, pack = excluded.pack, hf_repo = excluded.hf_repo, messages_json = excluded.messages_json, status = excluded.status, updated_at = excluded.updated_at """; cmd.Parameters.AddWithValue("$id", id); cmd.Parameters.AddWithValue("$ds", datasetId); cmd.Parameters.AddWithValue("$src", sample["source"]?.ToString() ?? "manual"); cmd.Parameters.AddWithValue("$chat", (object)sample["chat_id"]?.ToString() ?? DBNull.Value); cmd.Parameters.AddWithValue("$persona", (object)sample["persona"]?.ToString() ?? DBNull.Value); cmd.Parameters.AddWithValue("$pack", (object)sample["pack"]?.ToString() ?? DBNull.Value); cmd.Parameters.AddWithValue("$hf", (object)sample["hf_repo"]?.ToString() ?? DBNull.Value); cmd.Parameters.AddWithValue("$msg", messages.ToString(Newtonsoft.Json.Formatting.None)); cmd.Parameters.AddWithValue("$status", sample["status"]?.ToString() ?? "draft"); long created = sample["createdAt"]?.Value() ?? now; cmd.Parameters.AddWithValue("$c", created); cmd.Parameters.AddWithValue("$u", now); cmd.ExecuteNonQuery(); return new JObject { ["id"] = id, ["updatedAt"] = now }; } } public bool DeleteTrainSample(string id) { lock (_lock) { EnsureOpen(); using SqliteCommand cmd = _conn.CreateCommand(); cmd.CommandText = "DELETE FROM train_samples WHERE id = $id"; cmd.Parameters.AddWithValue("$id", id ?? ""); return cmd.ExecuteNonQuery() > 0; } } public int CountTrainSamples(string status = null) { lock (_lock) { EnsureOpen(); using SqliteCommand cmd = _conn.CreateCommand(); if (string.IsNullOrWhiteSpace(status) || string.Equals(status, "all", StringComparison.OrdinalIgnoreCase)) { cmd.CommandText = "SELECT COUNT(*) FROM train_samples"; } else { cmd.CommandText = "SELECT COUNT(*) FROM train_samples WHERE status = $s"; cmd.Parameters.AddWithValue("$s", status.Trim()); } return Convert.ToInt32(cmd.ExecuteScalar()); } } public JObject SaveTrainJob(JObject job) { lock (_lock) { EnsureOpen(); long now = DateTimeOffset.UtcNow.ToUnixTimeMilliseconds(); string id = job["id"]?.ToString()?.Trim(); if (string.IsNullOrWhiteSpace(id)) { id = $"tj_{now}_{Guid.NewGuid():N}"[..24]; } using SqliteCommand cmd = _conn.CreateCommand(); cmd.CommandText = """ INSERT INTO train_jobs(id, kind, status, config_json, base_model, output_name, log_path, progress_json, created_at, updated_at, finished_at) VALUES($id, $kind, $status, $cfg, $base, $out, $log, $prog, $c, $u, $f) ON CONFLICT(id) DO UPDATE SET status = excluded.status, config_json = excluded.config_json, log_path = excluded.log_path, progress_json = excluded.progress_json, updated_at = excluded.updated_at, finished_at = excluded.finished_at """; cmd.Parameters.AddWithValue("$id", id); cmd.Parameters.AddWithValue("$kind", job["kind"]?.ToString() ?? "qlora"); cmd.Parameters.AddWithValue("$status", job["status"]?.ToString() ?? "pending"); cmd.Parameters.AddWithValue("$cfg", job["config"]?.ToString(Newtonsoft.Json.Formatting.None) ?? job["config_json"]?.ToString() ?? "{}"); cmd.Parameters.AddWithValue("$base", (object)job["base_model"]?.ToString() ?? DBNull.Value); cmd.Parameters.AddWithValue("$out", (object)job["output_name"]?.ToString() ?? DBNull.Value); cmd.Parameters.AddWithValue("$log", (object)job["log_path"]?.ToString() ?? DBNull.Value); cmd.Parameters.AddWithValue("$prog", (object)job["progress"]?.ToString(Newtonsoft.Json.Formatting.None) ?? job["progress_json"]?.ToString() ?? DBNull.Value); cmd.Parameters.AddWithValue("$c", job["created_at"]?.Value() ?? job["createdAt"]?.Value() ?? now); cmd.Parameters.AddWithValue("$u", now); cmd.Parameters.AddWithValue("$f", (object)(job["finished_at"]?.Value() ?? job["finishedAt"]?.Value()) ?? DBNull.Value); cmd.ExecuteNonQuery(); return new JObject { ["id"] = id }; } } public JObject GetTrainJob(string id) { lock (_lock) { EnsureOpen(); using SqliteCommand cmd = _conn.CreateCommand(); cmd.CommandText = "SELECT id, kind, status, config_json, base_model, output_name, log_path, progress_json, created_at, updated_at, finished_at FROM train_jobs WHERE id = $id"; cmd.Parameters.AddWithValue("$id", id ?? ""); using SqliteDataReader r = cmd.ExecuteReader(); if (!r.Read()) { return null; } return new JObject { ["id"] = r.GetString(0), ["kind"] = r.GetString(1), ["status"] = r.GetString(2), ["config_json"] = r.IsDBNull(3) ? null : r.GetString(3), ["base_model"] = r.IsDBNull(4) ? null : r.GetString(4), ["output_name"] = r.IsDBNull(5) ? null : r.GetString(5), ["log_path"] = r.IsDBNull(6) ? null : r.GetString(6), ["progress_json"] = r.IsDBNull(7) ? null : r.GetString(7), ["created_at"] = r.GetInt64(8), ["updated_at"] = r.GetInt64(9), ["finished_at"] = r.IsDBNull(10) ? null : r.GetInt64(10), }; } } public JObject GetActiveTrainJob() { lock (_lock) { EnsureOpen(); using SqliteCommand cmd = _conn.CreateCommand(); cmd.CommandText = "SELECT id, kind, status, config_json, base_model, output_name, log_path, progress_json, created_at, updated_at, finished_at FROM train_jobs WHERE status IN ('pending','running') ORDER BY updated_at DESC LIMIT 1"; using SqliteDataReader r = cmd.ExecuteReader(); if (!r.Read()) { return null; } return new JObject { ["id"] = r.GetString(0), ["kind"] = r.GetString(1), ["status"] = r.GetString(2), ["config_json"] = r.IsDBNull(3) ? null : r.GetString(3), ["base_model"] = r.IsDBNull(4) ? null : r.GetString(4), ["output_name"] = r.IsDBNull(5) ? null : r.GetString(5), ["log_path"] = r.IsDBNull(6) ? null : r.GetString(6), ["progress_json"] = r.IsDBNull(7) ? null : r.GetString(7), ["created_at"] = r.GetInt64(8), ["updated_at"] = r.GetInt64(9), ["finished_at"] = r.IsDBNull(10) ? null : r.GetInt64(10), }; } } }