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