Added short-term memory

This commit is contained in:
2026-07-24 20:01:28 +03:00
parent 88ca4aecd2
commit 5cdf47fca8
6 changed files with 1244 additions and 42 deletions
Generated
+1095 -17
View File
File diff suppressed because it is too large Load Diff
+4 -1
View File
@@ -15,4 +15,7 @@ serde = { version = "1.0", features = ["derive"] }
serde_json = "1.0"
# Парсинг аргументов CLI (флаги --session-id и т.д.)
clap = { version = "4.5", features = ["derive"] }
clap = { version = "4.5", features = ["derive"] }
# Асинхронная работа с SQLite
sqlx = { version = "0.7", features = ["runtime-tokio-rustls", "sqlite"] }
BIN
View File
Binary file not shown.
Binary file not shown.
Binary file not shown.
+145 -24
View File
@@ -1,44 +1,129 @@
use clap::Parser;
use serde::{Deserialize, Serialize};
use serde_json::json;
use sqlx::sqlite::{SqliteConnectOptions, SqlitePoolOptions};
use sqlx::{Pool, Sqlite};
use std::io::{self, Write};
use std::str::FromStr;
use tokio::io::{AsyncBufReadExt, BufReader};
/// CLI-утилита для взаимодействия с LLM в рамках указанной сессии
#[derive(Parser, Debug)]
#[command(author, version, about, long_about = None)]
struct Args {
/// Идентификатор сессии/беседы (на будущее для VK peer_id)
#[arg(short, long, default_value = "default_session")]
session_id: String,
/// URL OpenAI-совместимого эндпоинта (vLLM / llama.cpp / Ollama)
#[arg(
short,
long,
default_value = "http://localhost"
default_value = "http://192.168.0.50:6969/v1/chat/completions"
)]
api_url: String,
/// Название модели (для Ollama/vLLM)
#[arg(short, long, default_value = "qwen2.5-coder")]
model: String,
/// Путь к файлу базы данных SQLite
#[arg(long, default_value = "sqlite:chat_history.db?mode=rwc")]
db_url: String,
}
/// Функция отправки Stateless-запроса к LLM
async fn send_llm_request(
#[derive(Serialize, Deserialize, Debug, Clone)]
struct ChatMessage {
role: String,
content: String,
}
/// Инициализация БД: создание таблицы, индексов и включение режима WAL
async fn init_db(db_url: &str) -> Result<Pool<Sqlite>, Box<dyn std::error::Error>> {
let options = SqliteConnectOptions::from_str(db_url)?
.create_if_missing(true)
.journal_mode(sqlx::sqlite::SqliteJournalMode::Wal); // WAL режим для высокой параллельности
let pool = SqlitePoolOptions::new()
.max_connections(10) // До 10 параллельных соединений к БД
.connect_with(options)
.await?;
// Создаем таблицу и индекс, если их нет
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS messages (
id INTEGER PRIMARY KEY AUTOINCREMENT,
session_id TEXT NOT NULL,
role TEXT NOT NULL,
content TEXT NOT NULL,
created_at DATETIME DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_messages_session_id_id
ON messages (session_id, id DESC);
"#,
)
.execute(&pool)
.await?;
Ok(pool)
}
/// Сохранение сообщения в историю
async fn save_message(
pool: &Pool<Sqlite>,
session_id: &str,
role: &str,
content: &str,
) -> Result<(), sqlx::Error> {
sqlx::query("INSERT INTO messages (session_id, role, content) VALUES (?, ?, ?)")
.bind(session_id)
.bind(role)
.bind(content)
.execute(pool)
.await?;
Ok(())
}
/// Выборка последних N (5) сообщений контекста в хронологическом порядке
async fn get_recent_history(
pool: &Pool<Sqlite>,
session_id: &str,
limit: i64,
) -> Result<Vec<ChatMessage>, sqlx::Error> {
let rows = sqlx::query_as::<_, (String, String)>(
r#"
SELECT role, content
FROM (
SELECT role, content, id
FROM messages
WHERE session_id = ?
ORDER BY id DESC
LIMIT ?
)
ORDER BY id ASC
"#,
)
.bind(session_id)
.bind(limit * 2)
.fetch_all(pool)
.await?;
let history = rows
.into_iter()
.map(|(role, content)| ChatMessage { role, content })
.collect();
Ok(history)
}
/// Отправка запроса с учетом накопленного контекста
async fn send_llm_request_with_history(
client: &reqwest::Client,
api_url: &str,
model: &str,
user_message: &str,
history: &[ChatMessage],
) -> Result<String, Box<dyn std::error::Error>> {
let payload = json!({
"model": model,
"messages": [
{
"role": "user",
"content": user_message
}
],
"messages": history, // Передаем весь сохраненный контекст
"temperature": 0.7
});
@@ -64,13 +149,23 @@ async fn main() {
let args = Args::parse();
let client = reqwest::Client::new();
// 1. Инициализируем БД
let db_pool = match init_db(&args.db_url).await {
Ok(pool) => pool,
Err(e) => {
eprintln!("Ошибка инициализации БД: {}", e);
return;
}
};
println!("==================================================");
println!(" LLM CLI Agent (Stage 1)");
println!(" LLM CLI Agent with Memory (Stage 2)");
println!(" Session ID: {}", args.session_id);
println!(" API URL: {}", args.api_url);
println!(" Model: {}", args.model);
println!(" Database: SQLite (WAL Mode)");
println!("==================================================");
println!("Введите сообщение и нажмите Enter. Для выхода наберите 'exit' или 'quit'.\n");
println!("Введите сообщение. Для выхода наберите 'exit' или 'quit'.\n");
let stdin = tokio::io::stdin();
let mut reader = BufReader::new(stdin);
@@ -83,7 +178,7 @@ async fn main() {
let bytes_read = reader.read_line(&mut input).await.unwrap_or(0);
if bytes_read == 0 {
break; // EOF (Ctrl+D)
break;
}
let trimmed_input = input.trim();
@@ -92,23 +187,49 @@ async fn main() {
continue;
}
if trimmed_input.eq_ignore_ascii_case("exit")
|| trimmed_input.eq_ignore_ascii_case("quit")
if trimmed_input.eq_ignore_ascii_case("exit") || trimmed_input.eq_ignore_ascii_case("quit")
{
println!("Завершение работы.");
break;
}
print!("[Ожидание ответа от LLM...]\r");
// 2. Сохраняем сообщение пользователя в БД
if let Err(e) = save_message(&db_pool, &args.session_id, "user", trimmed_input).await {
eprintln!("[Ошибка записи в БД]: {}", e);
continue;
}
// 3. Достаем последние 5 сообщений из БД (включая только что сохраненное)
let history = match get_recent_history(&db_pool, &args.session_id, 5).await {
Ok(h) => h,
Err(e) => {
eprintln!("[Ошибка чтения из БД]: {}", e);
continue;
}
};
print!(
"[Запрос к LLM с контекстом ({} сообщ.)...]\r",
history.len()
);
io::stdout().flush().unwrap();
match send_llm_request(&client, &args.api_url, &args.model, trimmed_input).await {
// 4. Отправляем контекст в LLM
match send_llm_request_with_history(&client, &args.api_url, &args.model, &history).await {
Ok(reply) => {
println!("\r[LLM]: {}\n", reply);
// \r - в начало, \x1B[2K - очистить всю текущую строку терминала
print!("\r\x1B[2K");
println!("[LLM]: {}\n", reply);
if let Err(e) = save_message(&db_pool, &args.session_id, "assistant", &reply).await
{
eprintln!("[Ошибка записи ответа в БД]: {}", e);
}
}
Err(e) => {
eprintln!("\r[Ошибка]: {}\n", e);
print!("\r\x1B[2K");
eprintln!("[Ошибка LLM]: {}\n", e);
}
}
}
}
}