test_server.py 1.1 KB

123456789101112131415161718192021222324252627282930
  1. import pytest
  2. from server import _extract_code, _read_only
  3. @pytest.mark.parametrize(
  4. ("language", "raw", "expected"),
  5. [
  6. ("SQL", "<think>hidden</think>```sql\nSELECT * FROM rescue_forces\n```", "SELECT * FROM rescue_forces"),
  7. ("SQL", "Here is the query:\nSELECT name FROM rescue_forces\nExplanation: done", "SELECT name FROM rescue_forces"),
  8. ("Cypher", "```cypher\nMATCH (m:Mission) RETURN m\n```", "MATCH (m:Mission) RETURN m"),
  9. ("Cypher", "思考\nMATCH (s:Stage) RETURN s\n说明:done", "MATCH (s:Stage) RETURN s"),
  10. ],
  11. )
  12. def test_extract_and_accept_read_only(language, raw, expected):
  13. assert _read_only(_extract_code(raw, language), language) == expected
  14. @pytest.mark.parametrize(
  15. ("language", "query"),
  16. [
  17. ("SQL", "DELETE FROM rescue_forces"),
  18. ("SQL", "SELECT 1; DROP TABLE rescue_forces"),
  19. ("Cypher", "MATCH (n) DELETE n"),
  20. ("Cypher", "MATCH (n) RETURN n; CREATE ()"),
  21. ],
  22. )
  23. def test_rejects_writes_and_multiple_statements(language, query):
  24. with pytest.raises(ValueError):
  25. _read_only(query, language)