//! Integration tests for RTX-NLG use rtx_nlg::*; use std::sync::Arc; #[tokio::test] async fn test_complete_generation_pipeline() -> Result<()> { // Create runtime let runtime = NlgRuntime::new()?; // Load a mock model runtime.load_model("test_model", "/tmp/mock").await?; // Test beam search generation let beam_config = GenerationConfig::beam_search() .num_beams(4) .max_length(50) .temperature(0.8); let mut generator = runtime.create_generator("test_model", beam_config)?; let output = generator.generate("The future of artificial intelligence")?; assert!(!output.text.is_empty()); assert!(output.metadata.tokens_per_second > 0.0); Ok(()) } #[tokio::test] async fn test_translation_pipeline() -> Result<()> { let config = translation::TranslationConfig::new("en", "es") .beam_size(4) .max_length(100); let translator = translation::Translator::new(config)?; let result = translator.translate("Hello, how are you today?")?; assert!(!result.translation.is_empty()); assert!(result.confidence > 0.0); assert!(result.metrics.translation_time_ms > 0.0); Ok(()) } #[tokio::test] async fn test_summarization_pipeline() -> Result<()> { let model = Arc::new(MockModelInterface::new("/tmp/mock")?); let config = summarization::SummarizationConfig::default(); let summarizer = summarization::Summarizer::new(model, config); let long_text = "This is a very long document that needs to be summarized. It contains multiple sentences with various pieces of information. The document discusses artificial intelligence, machine learning, and natural language processing. These technologies are becoming increasingly important in modern applications. They enable computers to understand and generate human language."; let summary = summarizer.summarize(long_text)?; assert!(!summary.text.is_empty()); assert!(summary.text.len() < long_text.len()); assert!(summary.metadata.quality_scores.fluency > 0.0); Ok(()) } #[tokio::test] async fn test_controllable_generation() -> Result<()> { let runtime = NlgRuntime::new()?; runtime.load_model("test_model", "/tmp/mock").await?; // Create controllable generation config let controllable_config = generation::controllable::ControllableConfig { sentiment_control: Some( generation::controllable::SentimentControl::new(0.8) .add_positive_word(100, 1.0) .add_positive_word(200, 0.8), ), control_strength: 0.7, ..Default::default() }; let gen_config = GenerationConfig { strategy: GenerationStrategy::Controllable(controllable_config), max_length: Some(50), ..Default::default() }; let mut generator = runtime.create_generator("test_model", gen_config)?; let output = generator.generate("Today is a wonderful day")?; assert!(!output.text.is_empty()); Ok(()) } #[tokio::test] async fn test_quality_filtering() -> Result<()> { let checker = quality::QualityChecker::new() .add_filter(Box::new(quality::RepetitionPenalty::default())) .add_filter(Box::new(quality::DiversityPromoter::default())) .add_filter(Box::new(quality::ToxicityFilter::default())); let good_text = "This is a well-written, diverse, and safe piece of text."; let bad_text = "same same same same hate hate hate violence violence"; let good_report = checker.check_quality(good_text)?; let bad_report = checker.check_quality(bad_text)?; assert!(good_report.passes_all); assert!(!bad_report.passes_all); assert!(good_report.overall_score > bad_report.overall_score); Ok(()) } #[tokio::test] async fn test_streaming_generation() -> Result<()> { let model = Arc::new(MockModelInterface::new("/tmp/mock")?); let config = GenerationConfig::default(); let generator = serving::StreamingGenerator::with_model(model, config)?; let stream = generator.generate_stream("Tell me a story").await?; use futures::StreamExt; let tokens: Vec<_> = stream.take(5).collect().await; assert!(!tokens.is_empty()); for token_result in tokens { assert!(token_result.is_ok()); let token = token_result?; assert!(!token.text.is_empty()); } Ok(()) } #[tokio::test] async fn test_batch_generation() -> Result<()> { let model = Arc::new(MockModelInterface::new("/tmp/mock")?); let config = GenerationConfig::default(); let generator = serving::BatchGenerator::with_model(model, config, 4)?; let prompts = vec![ "First prompt".to_string(), "Second prompt".to_string(), "Third prompt".to_string(), ]; let results = generator.generate_batch(prompts).await?; assert_eq!(results.len(), 3); for result in &results { assert!(!result.text.is_empty()); } Ok(()) } #[tokio::test] async fn test_conversation_management() -> Result<()> { let mut manager = serving::ConversationManager::new(); let conversation = manager.create_conversation("test_conv".to_string(), 10); assert_eq!(conversation.id, "test_conv"); manager.add_turn( "test_conv", serving::conversation::Role::User, "Hello".to_string(), )?; manager.add_turn( "test_conv", serving::conversation::Role::Assistant, "Hi there!".to_string(), )?; let context = manager.get_context("test_conv")?; assert!(context.contains("Hello")); assert!(context.contains("Hi there!")); Ok(()) } #[tokio::test] async fn test_template_rendering() -> Result<()> { let mut manager = serving::templates::TemplateManager::new(); let mut values = std::collections::HashMap::new(); values.insert("text".to_string(), "This is a test document.".to_string()); values.insert("max_sentences".to_string(), "2".to_string()); let rendered = manager.render_template("summarize", &values)?; assert!(rendered.contains("This is a test document.")); assert!(rendered.contains("2 sentences")); Ok(()) } #[test] fn test_optimization_kv_cache() -> Result<()> { let mut cache = optimization::KVCache::new(100)?; let device = rtx_tensor::Device::cuda(0).unwrap_or(Device::default()); let key = rtx_tensor::Tensor::zeros(&[1, 10, 64], &device)?; let value = rtx_tensor::Tensor::zeros(&[1, 10, 64], &device)?; cache.insert("test_key".to_string(), vec![key], vec![value])?; let stats = cache.usage_stats(); assert_eq!(stats.entries, 1); assert!(stats.size_mb > 0); Ok(()) } #[test] fn test_speculative_decoding() -> Result<()> { let decoder = optimization::SpeculativeDecoder::new(4)?; let target_model = MockModelInterface::new("/tmp/mock")?; let device = Device::default(); let input_ids = rtx_tensor::Tensor::from_slice(&[1f32, 2.0, 3.0], &[1, 3], &device)?; let result = decoder.speculative_step(&target_model, &input_ids, None)?; assert!(!result.accepted_tokens.is_empty()); assert!(result.acceptance_rate >= 0.0 && result.acceptance_rate <= 1.0); assert!(result.speedup_factor > 0.0); Ok(()) } #[tokio::test] async fn test_model_server() -> Result<()> { let config = serving::ServerConfig::default(); let server = serving::ModelServer::new(config); let model = Arc::new(MockModelInterface::new("/tmp/mock")?); server .register_model("test_model".to_string(), model) .await?; let request = serving::GenerationRequest { prompt: "Test prompt".to_string(), model_id: "test_model".to_string(), generation_config: GenerationConfig::default(), stream: None, conversation_id: None, }; let response = server.generate(request).await?; assert!(!response.text.is_empty()); assert_eq!(response.model_id, "test_model"); assert!(response.metadata.generation_time_ms > 0.0); Ok(()) } #[tokio::test] async fn test_end_to_end_workflow() -> Result<()> { // This test demonstrates a complete NLG workflow // 1. Initialize runtime let runtime = NlgRuntime::new()?; runtime.load_model("main_model", "/tmp/mock").await?; // 2. Create quality checker let quality_checker = quality::QualityChecker::new() .add_filter(Box::new(quality::RepetitionPenalty::default())) .add_filter(Box::new(quality::DiversityPromoter::default())) .add_filter(Box::new(quality::ToxicityFilter::default())); // 3. Generate text with quality control let config = GenerationConfig::nucleus_sampling() .top_p(0.9) .max_length(100) .temperature(0.8); let mut generator = runtime.create_generator("main_model", config)?; let mut attempts = 0; let max_attempts = 3; loop { let output = generator.generate("Write a creative story about")?; let quality_report = quality_checker.check_quality(&output.text)?; if quality_report.passes_all { // Quality check passed assert!(!output.text.is_empty()); assert!(quality_report.overall_score > 0.5); break; } attempts += 1; if attempts >= max_attempts { // Even if quality doesn't pass, test should not fail assert!(!output.text.is_empty()); break; } } Ok(()) }