use kcode_k1_chat_chatend::{BoxContent, ChatBox, Chatend};
pub use kcode_k1_chat_chatend::{BoxId, ToolCallId};
pub use kcode_k1_transaction_id::TxId;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::fs::{self, File, OpenOptions};
use std::io::{self, Seek, SeekFrom, Write};
use std::path::{Path, PathBuf};
pub type SessionId = [u8; 12];
pub const CODEC_VERSION: u32 = 1;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct EventRecord {
pub after_box_id: u64,
pub event_index: u64,
pub connected_box_id: u64,
pub handler: String,
pub data: Value,
}
impl EventRecord {
pub fn new(
after_box_id: u64,
event_index: u64,
connected_box_id: u64,
handler: String,
data: Value,
) -> Result<Self, Error> {
let value = Self {
after_box_id,
event_index,
connected_box_id,
handler,
data,
};
validate_event(&value)?;
Ok(value)
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum Record {
Box(ChatBox),
Event(EventRecord),
}
impl Record {
pub fn chat_box(value: ChatBox) -> Self {
Self::Box(value)
}
pub fn event(value: EventRecord) -> Self {
Self::Event(value)
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Batch {
pub version: u32,
pub session_id: SessionId,
pub predecessor: Option<Record>,
pub records: Vec<Record>,
}
impl Batch {
pub fn new(
session_id: SessionId,
predecessor: Option<Record>,
records: Vec<Record>,
) -> Result<Self, Error> {
let value = Self {
version: CODEC_VERSION,
session_id,
predecessor,
records,
};
validate_batch(&value)?;
Ok(value)
}
pub fn encode(&self) -> Result<Vec<u8>, Error> {
validate_batch(self)?;
Ok(serde_json::to_vec(&WireBatch::from(self))?)
}
pub fn decode(bytes: &[u8]) -> Result<Self, Error> {
let value = Self::try_from(serde_json::from_slice::<WireBatch>(bytes)?)?;
validate_batch(&value)?;
Ok(value)
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct SessionLog {
pub boxes: Vec<ChatBox>,
pub events: Vec<EventRecord>,
pub records: Vec<Record>,
}
#[derive(Debug)]
pub enum Error {
Io(io::Error),
Json(serde_json::Error),
Invalid(&'static str),
}
impl std::fmt::Display for Error {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Io(e) => write!(f, "I/O error: {e}"),
Self::Json(e) => write!(f, "JSON error: {e}"),
Self::Invalid(e) => f.write_str(e),
}
}
}
impl std::error::Error for Error {}
impl From<io::Error> for Error {
fn from(value: io::Error) -> Self {
Self::Io(value)
}
}
impl From<serde_json::Error> for Error {
fn from(value: serde_json::Error) -> Self {
Self::Json(value)
}
}
type Result<T, E = Error> = std::result::Result<T, E>;
pub struct Projection {
root: PathBuf,
sessions: PathBuf,
}
impl Projection {
pub fn new(root: impl AsRef<Path>) -> Result<Self> {
let root = root.as_ref().to_path_buf();
fs::create_dir_all(&root)?;
let sessions = root.join("sessions");
fs::create_dir_all(&sessions)?;
Ok(Self { root, sessions })
}
pub fn load(&self, session: SessionId) -> Result<SessionLog> {
let records = self
.read_strict(session)?
.into_iter()
.map(|value| value.1)
.collect::<Vec<_>>();
validate_session_records(session, &records)?;
Ok(to_log(records))
}
pub fn apply(&mut self, txid: TxId, batch: &Batch, reconcile_first: bool) -> Result<()> {
validate_batch(batch)?;
fs::create_dir_all(&self.sessions)?;
let path = self.path(batch.session_id);
if reconcile_first {
let bytes = match fs::read(&path) {
Ok(value) => value,
Err(error) if error.kind() == io::ErrorKind::NotFound => Vec::new(),
Err(error) => return Err(error.into()),
};
let end = locate_predecessor(&bytes, batch.predecessor.as_ref())?;
let mut combined = strict_prefix(&bytes[..end])?
.into_iter()
.map(|value| value.1)
.collect::<Vec<_>>();
combined.extend(batch.records.clone());
validate_session_records(batch.session_id, &combined)?;
let mut file = OpenOptions::new()
.create(true)
.read(true)
.write(true)
.truncate(false)
.open(&path)?;
file.set_len(end as u64)?;
file.seek(SeekFrom::Start(end as u64))?;
append_lines(&mut file, txid, &batch.records)?;
file.sync_all()?;
} else {
let current = self.read_strict(batch.session_id)?;
let mut combined = current
.iter()
.map(|value| value.1.clone())
.collect::<Vec<_>>();
if combined.last() != batch.predecessor.as_ref() {
return Err(Error::Invalid("predecessor mismatch"));
}
combined.extend(batch.records.clone());
validate_session_records(batch.session_id, &combined)?;
let mut file = OpenOptions::new().create(true).append(true).open(&path)?;
append_lines(&mut file, txid, &batch.records)?;
file.sync_all()?;
}
Ok(())
}
pub fn discard_all(&mut self) -> Result<()> {
match fs::remove_dir_all(&self.sessions) {
Ok(()) => {}
Err(error) if error.kind() == io::ErrorKind::NotFound => {}
Err(error) => return Err(error.into()),
}
File::open(&self.root)?.sync_all()?;
Ok(())
}
fn path(&self, session: SessionId) -> PathBuf {
self.sessions.join(format!("{}.jsonl", hex(session)))
}
fn read_strict(&self, session: SessionId) -> Result<Vec<(TxId, Record, usize)>> {
let bytes = match fs::read(self.path(session)) {
Ok(value) => value,
Err(error) if error.kind() == io::ErrorKind::NotFound => return Ok(Vec::new()),
Err(error) => return Err(error.into()),
};
strict_prefix(&bytes)
}
}
fn append_lines(file: &mut File, txid: TxId, records: &[Record]) -> Result<()> {
for record in records {
let mut bytes = serde_json::to_vec(&WireLine {
txid: hex(*txid.as_bytes()),
record: WireRecord::from(record),
})?;
bytes.push(b'\n');
file.write_all(&bytes)?;
}
Ok(())
}
fn strict_prefix(bytes: &[u8]) -> Result<Vec<(TxId, Record, usize)>> {
if !bytes.is_empty() && bytes.last() != Some(&b'\n') {
return Err(Error::Invalid("line is not newline terminated"));
}
let mut records = Vec::new();
let mut start = 0;
while start < bytes.len() {
let relative = bytes[start..]
.iter()
.position(|byte| *byte == b'\n')
.ok_or(Error::Invalid("line is not newline terminated"))?;
let end = start + relative + 1;
let wire: WireLine = serde_json::from_slice(&bytes[start..end - 1])?;
records.push((
TxId::from_bytes(parse_hex(&wire.txid)?),
Record::try_from(wire.record)?,
end,
));
start = end;
}
validate_records(
&records
.iter()
.map(|value| value.1.clone())
.collect::<Vec<_>>(),
)?;
Ok(records)
}
fn locate_predecessor(bytes: &[u8], wanted: Option<&Record>) -> Result<usize> {
let Some(wanted) = wanted else {
return Ok(0);
};
let mut start = 0;
let mut found = None;
let mut prefix = Vec::new();
while start < bytes.len() {
let Some(relative) = bytes[start..].iter().position(|byte| *byte == b'\n') else {
break;
};
let end = start + relative + 1;
let parsed = serde_json::from_slice::<WireLine>(&bytes[start..end - 1])
.map_err(Error::from)
.and_then(|wire| Record::try_from(wire.record));
let record = match parsed {
Ok(value) => value,
Err(_) if found.is_some() => break,
Err(error) => return Err(error),
};
if same_identity(&record, wanted) {
if &record != wanted {
return Err(Error::Invalid("predecessor identity mismatch"));
}
if found.is_some() {
return Err(Error::Invalid("duplicate predecessor identity"));
}
found = Some(end);
}
if found.is_none() {
prefix.push(record);
validate_records(&prefix)?;
}
start = end;
}
found.ok_or(Error::Invalid("predecessor absent"))
}
fn same_identity(left: &Record, right: &Record) -> bool {
match (left, right) {
(Record::Box(left), Record::Box(right)) => left.id() == right.id(),
(Record::Event(left), Record::Event(right)) => {
(left.after_box_id, left.event_index, left.connected_box_id)
== (
right.after_box_id,
right.event_index,
right.connected_box_id,
)
}
_ => false,
}
}
fn validate_batch(batch: &Batch) -> Result<()> {
if batch.version != CODEC_VERSION {
return Err(Error::Invalid("unsupported version"));
}
if batch.records.is_empty() {
return Err(Error::Invalid("empty batch"));
}
if let Some(record) = &batch.predecessor {
validate_record_session(batch.session_id, record)?;
}
for record in &batch.records {
validate_record_session(batch.session_id, record)?;
}
validate_suffix(batch.predecessor.as_ref(), &batch.records)?;
if batch.predecessor.is_none() {
validate_session_records(batch.session_id, &batch.records)?;
}
Ok(())
}
fn validate_suffix(predecessor: Option<&Record>, suffix: &[Record]) -> Result<()> {
let (mut latest, mut event_index) = match predecessor {
None => (0, 0),
Some(Record::Box(value)) => (value.id().get(), 0),
Some(Record::Event(value)) => {
validate_event(value)?;
(value.after_box_id, value.event_index)
}
};
for record in suffix {
match record {
Record::Box(value) => {
if value.id().get()
!= latest
.checked_add(1)
.ok_or(Error::Invalid("box ID overflow"))?
{
return Err(Error::Invalid("noncontiguous box ID"));
}
latest = value.id().get();
event_index = 0;
}
Record::Event(value) => {
validate_event(value)?;
if value.after_box_id != latest
|| value.event_index
!= event_index
.checked_add(1)
.ok_or(Error::Invalid("event index overflow"))?
|| value.connected_box_id > latest
{
return Err(Error::Invalid("invalid event order or association"));
}
event_index = value.event_index;
}
}
}
Ok(())
}
fn validate_event(value: &EventRecord) -> Result<()> {
if value.handler.is_empty() {
Err(Error::Invalid("empty event handler"))
} else {
Ok(())
}
}
fn validate_record_session(session: SessionId, record: &Record) -> Result<()> {
let id = match record {
Record::Box(value) => match value.content() {
BoxContent::KtoolCall { tool_call_id, .. }
| BoxContent::KtoolReturn { tool_call_id, .. } => Some(tool_call_id),
_ => None,
},
Record::Event(_) => None,
};
if id.is_some_and(|value| value.session() != session) {
Err(Error::Invalid("ToolCallId belongs to another session"))
} else {
Ok(())
}
}
fn validate_session_records(session: SessionId, records: &[Record]) -> Result<()> {
validate_records(records)?;
for record in records {
validate_record_session(session, record)?;
}
Ok(())
}
fn validate_records(records: &[Record]) -> Result<()> {
validate_suffix(None, records)?;
let boxes = records
.iter()
.filter_map(|record| match record {
Record::Box(value) => Some(value.clone()),
Record::Event(_) => None,
})
.collect();
Chatend::recover(boxes).map_err(|_| Error::Invalid("invalid Chatend recovery"))?;
Ok(())
}
fn to_log(records: Vec<Record>) -> SessionLog {
let boxes = records
.iter()
.filter_map(|record| match record {
Record::Box(value) => Some(value.clone()),
Record::Event(_) => None,
})
.collect();
let events = records
.iter()
.filter_map(|record| match record {
Record::Event(value) => Some(value.clone()),
Record::Box(_) => None,
})
.collect();
SessionLog {
boxes,
events,
records,
}
}
fn hex(bytes: [u8; 12]) -> String {
bytes.iter().map(|value| format!("{value:02x}")).collect()
}
fn parse_hex(value: &str) -> Result<[u8; 12]> {
if value.len() != 24
|| !value
.bytes()
.all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte))
{
return Err(Error::Invalid("invalid lowercase 24-hex value"));
}
let mut output = [0; 12];
for (index, slot) in output.iter_mut().enumerate() {
*slot = u8::from_str_radix(&value[index * 2..index * 2 + 2], 16)
.map_err(|_| Error::Invalid("invalid hex"))?;
}
Ok(output)
}
#[derive(Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
struct WireBatch {
version: u32,
session_id: String,
predecessor: Option<WireRecord>,
records: Vec<WireRecord>,
}
impl From<&Batch> for WireBatch {
fn from(value: &Batch) -> Self {
Self {
version: value.version,
session_id: hex(value.session_id),
predecessor: value.predecessor.as_ref().map(WireRecord::from),
records: value.records.iter().map(WireRecord::from).collect(),
}
}
}
impl TryFrom<WireBatch> for Batch {
type Error = Error;
fn try_from(value: WireBatch) -> Result<Self> {
Ok(Self {
version: value.version,
session_id: parse_hex(&value.session_id)?,
predecessor: value.predecessor.map(Record::try_from).transpose()?,
records: value
.records
.into_iter()
.map(Record::try_from)
.collect::<Result<_>>()?,
})
}
}
#[derive(Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
struct WireLine {
txid: String,
record: WireRecord,
}
#[derive(Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
enum WireRecord {
Box {
id: u64,
content: WireContent,
},
Event {
after_box_id: u64,
event_index: u64,
connected_box_id: u64,
handler: String,
data: Value,
},
}
impl From<&Record> for WireRecord {
fn from(value: &Record) -> Self {
match value {
Record::Box(value) => Self::Box {
id: value.id().get(),
content: WireContent::from(value.content()),
},
Record::Event(value) => Self::Event {
after_box_id: value.after_box_id,
event_index: value.event_index,
connected_box_id: value.connected_box_id,
handler: value.handler.clone(),
data: value.data.clone(),
},
}
}
}
impl TryFrom<WireRecord> for Record {
type Error = Error;
fn try_from(value: WireRecord) -> Result<Self> {
Ok(match value {
WireRecord::Box { id, content } => {
Self::Box(ChatBox::new(BoxId::new(id), BoxContent::try_from(content)?))
}
WireRecord::Event {
after_box_id,
event_index,
connected_box_id,
handler,
data,
} => Self::Event(EventRecord::new(
after_box_id,
event_index,
connected_box_id,
handler,
data,
)?),
})
}
}
#[derive(Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case", deny_unknown_fields)]
enum WireContent {
System {
text: String,
},
User {
text: String,
},
Kennedy {
text: String,
},
Attachment,
KtoolCall {
session: String,
sequence: u64,
name: String,
arguments: String,
},
KtoolReturn {
session: String,
sequence: u64,
originating_call: u64,
result: WireResult,
},
}
impl From<&BoxContent> for WireContent {
fn from(value: &BoxContent) -> Self {
match value {
BoxContent::System(text) => Self::System { text: text.clone() },
BoxContent::User(text) => Self::User { text: text.clone() },
BoxContent::Kennedy { text } => Self::Kennedy { text: text.clone() },
BoxContent::Attachment => Self::Attachment,
BoxContent::KtoolCall {
tool_call_id,
name,
arguments,
} => Self::KtoolCall {
session: hex(tool_call_id.session()),
sequence: tool_call_id.sequence(),
name: name.clone(),
arguments: arguments.clone(),
},
BoxContent::KtoolReturn {
tool_call_id,
originating_call,
result,
} => Self::KtoolReturn {
session: hex(tool_call_id.session()),
sequence: tool_call_id.sequence(),
originating_call: originating_call.get(),
result: WireResult::from(result),
},
}
}
}
impl TryFrom<WireContent> for BoxContent {
type Error = Error;
fn try_from(value: WireContent) -> Result<Self> {
Ok(match value {
WireContent::System { text } => Self::System(text),
WireContent::User { text } => Self::User(text),
WireContent::Kennedy { text } => Self::Kennedy { text },
WireContent::Attachment => Self::Attachment,
WireContent::KtoolCall {
session,
sequence,
name,
arguments,
} => Self::KtoolCall {
tool_call_id: ToolCallId::new(parse_hex(&session)?, sequence),
name,
arguments,
},
WireContent::KtoolReturn {
session,
sequence,
originating_call,
result,
} => Self::KtoolReturn {
tool_call_id: ToolCallId::new(parse_hex(&session)?, sequence),
originating_call: BoxId::new(originating_call),
result: result.into(),
},
})
}
}
#[derive(Serialize, Deserialize)]
#[serde(tag = "status", rename_all = "snake_case", deny_unknown_fields)]
enum WireResult {
Ok { value: String },
Err { value: String },
}
impl From<&std::result::Result<String, String>> for WireResult {
fn from(value: &std::result::Result<String, String>) -> Self {
match value {
Ok(value) => Self::Ok {
value: value.clone(),
},
Err(value) => Self::Err {
value: value.clone(),
},
}
}
}
impl From<WireResult> for std::result::Result<String, String> {
fn from(value: WireResult) -> Self {
match value {
WireResult::Ok { value } => Ok(value),
WireResult::Err { value } => Err(value),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
use std::sync::atomic::{AtomicU64, Ordering};
fn boxed(id: u64, content: BoxContent) -> Record {
Record::Box(ChatBox::new(BoxId::new(id), content))
}
fn event(after: u64, index: u64) -> Record {
Record::Event(
EventRecord::new(
after,
index,
after,
"handler".into(),
json!({"z":[true,null,{"a":1}]}),
)
.unwrap(),
)
}
fn directory() -> PathBuf {
static NEXT: AtomicU64 = AtomicU64::new(0);
let path = std::env::temp_dir().join(format!(
"k1-persistence-{}-{}",
std::process::id(),
NEXT.fetch_add(1, Ordering::Relaxed)
));
let _ = fs::remove_dir_all(&path);
path
}
fn tx(value: u8) -> TxId {
TxId::from_bytes([value; 12])
}
#[test]
fn codec_order_and_session_validation() {
let tool = ToolCallId::new([1; 12], 7);
let records = vec![
event(0, 1),
boxed(1, BoxContent::System("s".into())),
boxed(2, BoxContent::User("u".into())),
boxed(3, BoxContent::Kennedy { text: "k".into() }),
boxed(4, BoxContent::Attachment),
boxed(
5,
BoxContent::KtoolCall {
tool_call_id: tool,
name: "n".into(),
arguments: "{}".into(),
},
),
boxed(
6,
BoxContent::KtoolReturn {
tool_call_id: tool,
originating_call: BoxId::new(5),
result: Err("e".into()),
},
),
event(6, 1),
];
let batch = Batch::new([1; 12], None, records).unwrap();
assert_eq!(Batch::decode(&batch.encode().unwrap()).unwrap(), batch);
let wrong = boxed(
1,
BoxContent::KtoolCall {
tool_call_id: ToolCallId::new([2; 12], 1),
name: "n".into(),
arguments: "{}".into(),
},
);
assert!(Batch::new([1; 12], None, vec![wrong]).is_err());
assert!(
Batch::new(
[2; 12],
None,
vec![boxed(
1,
BoxContent::KtoolCall {
tool_call_id: ToolCallId::new([2; 12], 1),
name: "n".into(),
arguments: "{}".into(),
},
)],
)
.is_ok()
);
assert!(Batch::new([0; 12], None, vec![boxed(2, BoxContent::Attachment)]).is_err());
assert!(EventRecord::new(0, 1, 0, String::new(), json!(null)).is_err());
}
#[test]
fn append_load_reconcile_and_discard() {
let root = directory();
let mut projection = Projection::new(&root).unwrap();
let first = Batch::new(
[4; 12],
None,
vec![boxed(1, BoxContent::System("a".into()))],
)
.unwrap();
projection.apply(tx(1), &first, false).unwrap();
let next = Batch::new(
[4; 12],
Some(first.records[0].clone()),
vec![event(1, 1), boxed(2, BoxContent::User("b".into()))],
)
.unwrap();
projection.apply(tx(2), &next, false).unwrap();
projection.apply(tx(2), &next, true).unwrap();
assert_eq!(projection.load([4; 12]).unwrap().records.len(), 3);
let mut file = OpenOptions::new()
.append(true)
.open(projection.path([4; 12]))
.unwrap();
file.write_all(b"{partial").unwrap();
projection.apply(tx(2), &next, true).unwrap();
assert_eq!(projection.load([4; 12]).unwrap().records.len(), 3);
assert!(projection.load([5; 12]).unwrap().records.is_empty());
projection.discard_all().unwrap();
assert!(projection.load([4; 12]).unwrap().records.is_empty());
fs::remove_dir_all(root).unwrap();
}
#[test]
fn malformed_lines_and_predecessors_fail_closed() {
let root = directory();
let mut projection = Projection::new(&root).unwrap();
fs::write(projection.path([1; 12]), b"{}\n").unwrap();
assert!(projection.load([1; 12]).is_err());
fs::write(projection.path([2; 12]), b"{}").unwrap();
assert!(projection.load([2; 12]).is_err());
let missing = Batch::new(
[3; 12],
Some(boxed(9, BoxContent::Attachment)),
vec![boxed(10, BoxContent::Attachment)],
)
.unwrap();
assert!(projection.apply(tx(3), &missing, true).is_err());
fs::remove_dir_all(root).unwrap();
}
}