kcode-k1-chat-chatend-testkit 0.1.0

Concrete downstream conformance assertions for K1 Chatend
Documentation
#![forbid(unsafe_code)]

use kcode_k1_chat_chatend::{
    BoxId, ChatBox, Chatend, ProviderCall, RecoveryError, TOOL_CALL_HIDDEN_TYPE, TOOL_CALL_TYPE,
    TOOL_MESSAGE_TYPE, TOOL_RESULT_TYPE, ToolCallId, ToolMessageMetadata, ToolResultV2Metadata,
    TransitionError, USER_MESSAGE_TYPE, tool_call_box, tool_message_box, tool_result_box,
    tool_result_v2_box,
};

pub fn assert_chatend_conformance() {
    assert_active_lifecycle();
    assert_done_order();
    assert_result_without_messages();
    assert_recovery_rejections();
    assert_convention_boundaries();
}

fn assert_active_lifecycle() {
    let call_id = tool_id(1, 7);
    let mut chat = Chatend::new();
    assert_eq!(
        chat.accept_user("question".into()).unwrap(),
        Some(BoxId::new(1))
    );
    chat.start_round().unwrap();
    let dispatched = chat
        .append_stage("calling".into(), vec![provider_call(call_id)])
        .unwrap();
    assert_eq!(dispatched[0].call_box_id, BoxId::new(3));
    assert_eq!(
        chat.accept_tool_message(call_id, "first".into()).unwrap(),
        None
    );
    assert_eq!(
        chat.accept_tool_message(call_id, "second".into()).unwrap(),
        None
    );
    assert_eq!(
        chat.accept_async_return_v2(
            call_id,
            Ok("answer".into()),
            "k1.web-search-result/v1".into(),
            r#"{"sources":[]}"#.into(),
        )
        .unwrap(),
        None
    );

    let arrivals = chat.flush_active_arrivals().unwrap();
    assert_eq!(
        arrivals.iter().map(ChatBox::box_type).collect::<Vec<_>>(),
        vec![TOOL_MESSAGE_TYPE, TOOL_MESSAGE_TYPE, TOOL_RESULT_TYPE]
    );
    for (offset, value) in arrivals.iter().enumerate() {
        assert_eq!(value.id(), BoxId::new(4 + offset as u64));
    }
    for (offset, value) in arrivals[..2].iter().enumerate() {
        let metadata = value.tool_message_metadata().unwrap().unwrap();
        assert_eq!(metadata.originating_call, BoxId::new(3));
        assert_eq!(metadata.message_index, 1 + offset as u64);
    }
    let result = arrivals[2].tool_result_v2_metadata().unwrap().unwrap();
    assert_eq!(result.originating_call, BoxId::new(3));
    assert_eq!(result.metadata_type, "k1.web-search-result/v1");
    assert_eq!(
        chat.accept_tool_message(call_id, "late".into()),
        Err(TransitionError::ToolMessageAfterResult)
    );
    assert_eq!(
        chat.accept_async_return(call_id, Ok("duplicate".into())),
        Err(TransitionError::DuplicateReturn)
    );
    Chatend::recover(chat.boxes().to_vec()).unwrap();
}

fn assert_done_order() {
    let call_id = tool_id(2, 1);
    let mut chat = Chatend::new();
    chat.start_round().unwrap();
    chat.append_stage(String::new(), vec![provider_call(call_id)])
        .unwrap();
    chat.accept_tool_message(call_id, "interim".into()).unwrap();
    chat.accept_async_return(call_id, Ok("result".into()))
        .unwrap();
    let appended = chat.done("provider-final".into()).unwrap();
    assert_eq!(
        appended.iter().map(ChatBox::box_type).collect::<Vec<_>>(),
        vec!["Agent Message", TOOL_MESSAGE_TYPE, TOOL_RESULT_TYPE]
    );
    assert!(!chat.round_active());
}

fn assert_result_without_messages() {
    let call_id = tool_id(3, 1);
    let mut chat = Chatend::new();
    chat.start_round().unwrap();
    chat.append_stage(String::new(), vec![provider_call(call_id)])
        .unwrap();
    chat.accept_async_return(call_id, Ok("done".into()))
        .unwrap();
    let arrivals = chat.flush_active_arrivals().unwrap();
    assert_eq!(arrivals.len(), 1);
    assert_eq!(arrivals[0].box_type(), TOOL_RESULT_TYPE);
    Chatend::recover(chat.boxes().to_vec()).unwrap();
}

fn assert_recovery_rejections() {
    let call_id = tool_id(4, 9);
    let call = with_id(tool_call_box(&provider_call(call_id)), 1);

    let gap = with_id(
        tool_message_box(&ToolMessageMetadata {
            tool_call_id: call_id,
            originating_call: BoxId::new(1),
            message_index: 2,
            message: "gap".into(),
        })
        .unwrap(),
        2,
    );
    assert_recovery_error(
        vec![call.clone(), gap],
        RecoveryError::NonConsecutiveToolMessage,
    );

    let wrong_origin = with_id(
        tool_message_box(&ToolMessageMetadata {
            tool_call_id: call_id,
            originating_call: BoxId::new(99),
            message_index: 1,
            message: "wrong".into(),
        })
        .unwrap(),
        2,
    );
    assert_recovery_error(
        vec![call.clone(), wrong_origin],
        RecoveryError::WrongOriginatingCall,
    );

    let orphan = with_id(
        tool_message_box(&ToolMessageMetadata {
            tool_call_id: tool_id(8, 1),
            originating_call: BoxId::new(1),
            message_index: 1,
            message: "orphan".into(),
        })
        .unwrap(),
        2,
    );
    assert_recovery_error(vec![call.clone(), orphan], RecoveryError::UnknownToolCall);

    let result = with_id(tool_result_box(call_id, BoxId::new(1), Ok("one".into())), 2);
    let duplicate = with_id(
        tool_result_v2_box(&ToolResultV2Metadata {
            tool_call_id: call_id,
            originating_call: BoxId::new(1),
            result: Ok("two".into()),
            metadata_type: "test".into(),
            metadata_contents: "{}".into(),
        }),
        3,
    );
    assert_recovery_error(
        vec![call.clone(), result.clone(), duplicate],
        RecoveryError::DuplicateReturn,
    );

    let late = with_id(
        tool_message_box(&ToolMessageMetadata {
            tool_call_id: call_id,
            originating_call: BoxId::new(1),
            message_index: 1,
            message: "late".into(),
        })
        .unwrap(),
        3,
    );
    assert_recovery_error(
        vec![call.clone(), result, late],
        RecoveryError::ToolMessageAfterResult,
    );

    assert_recovery_error(vec![with_id(call, 2)], RecoveryError::NonContiguousBoxId);
}

fn assert_convention_boundaries() {
    let malformed = ChatBox::new(
        BoxId::new(1),
        USER_MESSAGE_TYPE.into(),
        String::new(),
        TOOL_CALL_HIDDEN_TYPE.into(),
        String::new(),
    );
    assert_recovery_error(vec![malformed], RecoveryError::MalformedToolConvention);

    let unknown = ChatBox::new(
        BoxId::new(1),
        TOOL_CALL_TYPE.into(),
        "future".into(),
        "future/tool-call".into(),
        "opaque".into(),
    );
    let recovered = Chatend::recover(vec![unknown.clone()]).unwrap();
    assert_eq!(recovered.boxes(), &[unknown]);
}

fn provider_call(tool_call_id: ToolCallId) -> ProviderCall {
    ProviderCall {
        tool_call_id,
        name: "WebSearch".into(),
        arguments: "{}".into(),
    }
}

fn tool_id(byte: u8, sequence: u64) -> ToolCallId {
    ToolCallId::new([byte; 12], sequence)
}

fn with_id(value: ChatBox, id: u64) -> ChatBox {
    ChatBox::new(
        BoxId::new(id),
        value.box_type().to_owned(),
        value.contents().to_owned(),
        value.hidden_type().to_owned(),
        value.hidden_contents().to_owned(),
    )
}

fn assert_recovery_error(boxes: Vec<ChatBox>, expected: RecoveryError) {
    assert!(matches!(Chatend::recover(boxes), Err(error) if error == expected));
}

#[cfg(test)]
mod tests {
    #[test]
    fn complete_chatend_conformance() {
        super::assert_chatend_conformance();
    }
}