import unittest
from langchain_core.documents import Document
from example import Retriever,build_pipeline,load_fixture
class RagTests(unittest.TestCase):
    def setUp(self): self.r=Retriever().replace(load_fixture())
    def test_known_source_and_prompt(self):
        docs=self.r.retrieve('refund window')
        self.assertEqual([d.metadata['source'] for d in docs],['refund-policy'])
        text=build_pipeline(self.r).invoke('refund window').to_string()
        self.assertIn('[refund-policy]',text); self.assertIn('14 days',text)
    def test_no_match_and_access_scope(self):
        self.assertEqual(self.r.retrieve('orbital velocity'),[])
        self.assertEqual(self.r.retrieve('payroll secret'),[])
        self.assertEqual(len(self.r.retrieve('payroll secret',('staff',))),1)
        self.assertIn('(no matching context)',build_pipeline(self.r).invoke('orbital velocity').to_string())
    def test_replace_idempotence_and_deletion(self):
        ids=[d.metadata['chunk_id'] for d in self.r.chunks]
        self.r.replace(load_fixture())
        self.assertEqual([d.metadata['chunk_id'] for d in self.r.chunks],ids)
        self.r.replace(load_fixture()[1:])
        self.assertEqual(self.r.retrieve('refund window'),[])
    def test_changed_content_gets_changed_id(self):
        docs=load_fixture(); old=self.r.chunks[0].metadata['chunk_id']
        docs[0].page_content='Refund window is 30 days.'; self.r.replace(docs)
        self.assertNotEqual(self.r.chunks[0].metadata['chunk_id'],old)
    def test_splitter_keeps_source_on_long_text(self):
        doc=Document(page_content='refund policy sentence. '*30,metadata={'source':'long','scope':'public'})
        self.r.replace([doc]); self.assertGreater(len(self.r.chunks),1)
        self.assertTrue(all(d.metadata['source']=='long' for d in self.r.chunks))
        self.assertEqual(len({d.metadata['chunk_id'] for d in self.r.chunks}),len(self.r.chunks))
    def test_invalid_inputs(self):
        docs=load_fixture()
        with self.assertRaises(ValueError): self.r.replace([docs[0],docs[0]])
        with self.assertRaises(ValueError): self.r.retrieve('')
        with self.assertRaises(ValueError): self.r.retrieve('refund',('unknown',))
        with self.assertRaises(ValueError): self.r.retrieve('refund',k=0)
if __name__=='__main__': unittest.main()
