Ship Assistent 0.12.1: training tab, dataset pipeline, and heard RAG.
Restructure UI with app-level tabs and chat history drawer; add dataset curation, HF import, Modelfile/QLoRA hooks, and link approved samples to the agent immediately via heard vector memory without waiting for fine-tuning. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -0,0 +1,345 @@
|
||||
using System;
|
||||
using System.Collections.Generic;
|
||||
using System.Linq;
|
||||
using Microsoft.Data.Sqlite;
|
||||
using Newtonsoft.Json.Linq;
|
||||
using SwarmUI.Utils;
|
||||
|
||||
namespace Mrleo1nid.SwarmAssistent;
|
||||
|
||||
/// <summary>Training datasets, samples, and job metadata in assistent.sqlite.</summary>
|
||||
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<JObject> ListTrainSamples(string status = null, string persona = null, string datasetId = null, int limit = 200)
|
||||
{
|
||||
lock (_lock)
|
||||
{
|
||||
EnsureOpen();
|
||||
int take = Math.Clamp(limit, 1, 2000);
|
||||
List<string> 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<JObject> 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<long?>() ?? 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<long?>() ?? job["createdAt"]?.Value<long?>() ?? now);
|
||||
cmd.Parameters.AddWithValue("$u", now);
|
||||
cmd.Parameters.AddWithValue("$f", (object)(job["finished_at"]?.Value<long?>() ?? job["finishedAt"]?.Value<long?>()) ?? 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),
|
||||
};
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user