Add missing WebSocket/HttpClient usings, stop using static on Config/FilePath helpers, and copy Microsoft.Data.Sqlite next to the extension dll. Co-authored-by: Cursor <cursoragent@cursor.com>
346 lines
15 KiB
C#
346 lines
15 KiB
C#
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),
|
|
};
|
|
}
|
|
}
|
|
}
|