new feature: Integration with kbdb Autonomous

This commit is contained in:
T3782834
2026-08-26 10:15:51 -03:00
parent ac18d68eaf
commit faf5ca55ba
405 changed files with 2368 additions and 138 deletions

View File

@@ -0,0 +1,160 @@
import base64
import json
from types import SimpleNamespace
import pytest
from agent_framework.checkpoints.langgraph_saver import (
RepositoryCheckpointSaver,
_TYPED_SERDE_MARKER,
_decode_checkpoint_value,
_encode_checkpoint_value,
_encode_pending_write_value,
)
Interrupt = type("Interrupt", (), {})
Interrupt.__module__ = "langgraph.types"
class FakeSerde:
def dumps_typed(self, value):
assert type(value).__module__ == "langgraph.types"
assert type(value).__qualname__ == "Interrupt"
return "msgpack", b"interrupt-payload"
def loads_typed(self, typed):
type_name, payload = typed
assert type_name == "msgpack"
assert payload == b"interrupt-payload"
return Interrupt()
class MemoryRepo:
def __init__(self):
self.value = None
async def put(self, thread_id, value):
self.value = value
async def get_latest(self, thread_id):
return self.value
def test_checkpoint_interrupt_only_in_special_channel_round_trips():
serde = FakeSerde()
raw_interrupt = Interrupt()
checkpoint = {
"id": "cp1",
"channel_values": {
"__root__": {
"__interrupt__": [raw_interrupt],
"business": {"ok": True},
},
"ordinary": {"value": 1},
},
}
encoded = _encode_checkpoint_value(serde, checkpoint)
leaf = encoded["channel_values"]["__root__"]["__interrupt__"][0]
assert leaf[_TYPED_SERDE_MARKER] is True
assert leaf["type"] == "msgpack"
assert base64.b64decode(leaf["data"]) == b"interrupt-payload"
assert encoded["channel_values"]["ordinary"] == {"value": 1}
decoded = _decode_checkpoint_value(serde, encoded)
assert type(decoded["channel_values"]["__root__"]["__interrupt__"][0]).__module__ == "langgraph.types"
assert decoded["channel_values"]["ordinary"] == {"value": 1}
def test_checkpoint_interrupt_outside_special_channel_is_rejected():
serde = FakeSerde()
checkpoint = {
"id": "cp1",
"channel_values": {"ordinary": [Interrupt()]},
}
with pytest.raises(TypeError, match="não serializável"):
_encode_checkpoint_value(serde, checkpoint)
def test_checkpoint_unknown_object_inside_interrupt_channel_is_rejected():
serde = FakeSerde()
class Unknown:
pass
checkpoint = {
"id": "cp1",
"channel_values": {"__root__": {"__interrupt__": [Unknown()]}},
}
with pytest.raises(TypeError, match="não serializável"):
_encode_checkpoint_value(serde, checkpoint)
def test_plain_checkpoint_keeps_plain_json_shape():
serde = FakeSerde()
checkpoint = {
"id": "cp1",
"channel_values": {"x": {"nested": [1, "a", True, None]}},
"channel_versions": {"x": "1"},
}
encoded = _encode_checkpoint_value(serde, checkpoint)
assert encoded == checkpoint
assert _TYPED_SERDE_MARKER not in json.dumps(encoded)
def test_pending_write_interrupt_requires_interrupt_branch():
serde = FakeSerde()
encoded = _encode_pending_write_value(
serde,
{"__interrupt__": [Interrupt()]},
path="$.pending_writes[t].__root__",
)
assert encoded["__interrupt__"][0][_TYPED_SERDE_MARKER] is True
with pytest.raises(TypeError, match="não serializável"):
_encode_pending_write_value(
serde,
{"ordinary": [Interrupt()]},
path="$.pending_writes[t].__root__",
)
def test_aput_encodes_only_checkpoint_interrupt_and_keeps_other_sections_strict():
repo = MemoryRepo()
settings = SimpleNamespace(CHECKPOINT_REPOSITORY_PROVIDER="memory")
saver = RepositoryCheckpointSaver(settings, repository=repo)
saver.serde = FakeSerde()
checkpoint = {
"id": "cp1",
"channel_values": {"__root__": {"__interrupt__": [Interrupt()]}},
}
import asyncio
asyncio.run(saver.aput(
{"configurable": {"thread_id": "t1"}},
checkpoint,
{"source": "test"},
{"x": 1},
))
assert repo.value["metadata"] == {"source": "test"}
assert repo.value["new_versions"] == {"x": 1}
assert repo.value["checkpoint"]["channel_values"]["__root__"]["__interrupt__"][0][_TYPED_SERDE_MARKER] is True
def test_aput_does_not_enable_typed_serde_for_metadata():
repo = MemoryRepo()
settings = SimpleNamespace(CHECKPOINT_REPOSITORY_PROVIDER="memory")
saver = RepositoryCheckpointSaver(settings, repository=repo)
saver.serde = FakeSerde()
import asyncio
with pytest.raises(TypeError, match="metadata"):
asyncio.run(saver.aput(
{"configurable": {"thread_id": "t1"}},
{"id": "cp1", "channel_values": {"x": 1}},
{"bad": Interrupt()},
{},
))

View File

@@ -1,6 +1,7 @@
from __future__ import annotations
import asyncio
import pytest
from types import SimpleNamespace
from agent_framework.checkpoints.langgraph_saver import (
@@ -193,3 +194,104 @@ def test_checkpoint_serialization_never_silently_stringifies_unknown_objects() -
with pytest.raises(TypeError, match="Checkpoint contém valor não serializável"):
_strict_json_value({"channel_values": {"bad": RuntimeLike()}}, path="$.checkpoint")
def test_typed_serializer_bridge_is_limited_to_langgraph_interrupt_pending_write() -> None:
import base64
import json
Interrupt = type("Interrupt", (), {"__module__": "langgraph.types"})
class FakeSerde:
def dumps_typed(self, value):
assert type(value).__module__ == "langgraph.types"
assert type(value).__qualname__ == "Interrupt"
return ("fake_interrupt", json.dumps({"value": value.value}).encode("utf-8"))
def loads_typed(self, payload):
type_name, raw = payload
assert type_name == "fake_interrupt"
obj = Interrupt()
obj.value = json.loads(raw.decode("utf-8"))["value"]
return obj
interrupt = Interrupt()
interrupt.value = "pause"
repo = _Repo()
saver = RepositoryCheckpointSaver(SimpleNamespace(), repository=repo)
saver.serde = FakeSerde()
config = {"configurable": {"thread_id": "typed-thread"}}
asyncio.run(saver.aput(config, {"id": "cp-typed", "channel_values": {}}, {}, {}))
# Mirrors the real LangGraph shape observed in the runtime log:
# pending_writes -> __root__ -> __interrupt__ -> [Interrupt(...)]
asyncio.run(saver.aput_writes(
config,
[("__root__", {"__interrupt__": [interrupt]})],
"task-typed",
))
persisted = repo.saved[1]
stored = persisted["pending_writes"][0]["value"]["__interrupt__"][0]
assert stored["__agent_framework_langgraph_typed__"] is True
assert stored["type"] == "fake_interrupt"
assert base64.b64decode(stored["data"]).decode("utf-8") == '{"value": "pause"}'
# Critical compatibility assertion: checkpoint/metadata/new_versions remain plain JSON.
assert persisted["checkpoint"] == {"id": "cp-typed", "channel_values": {}, "channel_versions": {}, "versions_seen": {}}
assert persisted["metadata"] == {}
restored = saver._make_tuple(persisted)
pending = restored.pending_writes if hasattr(restored, "pending_writes") else restored["pending_writes"]
assert pending[0][0] == "task-typed"
assert pending[0][1] == "__root__"
restored_interrupt = pending[0][2]["__interrupt__"][0]
assert type(restored_interrupt).__module__ == "langgraph.types"
assert type(restored_interrupt).__qualname__ == "Interrupt"
assert restored_interrupt.value == "pause"
def test_non_interrupt_non_json_pending_write_still_fails_loudly() -> None:
class RuntimeLike:
pass
repo = _Repo()
saver = RepositoryCheckpointSaver(SimpleNamespace(), repository=repo)
config = {"configurable": {"thread_id": "bad-pending-thread"}}
asyncio.run(saver.aput(config, {"id": "cp-bad", "channel_values": {}}, {}, {}))
with pytest.raises(TypeError, match="Checkpoint contém valor não serializável"):
asyncio.run(saver.aput_writes(config, [("__root__", {"bad": RuntimeLike()})], "task-bad"))
def test_non_json_checkpoint_does_not_fall_back_to_typed_serde() -> None:
class RuntimeLike:
pass
class ExplodingSerde:
def dumps_typed(self, value):
raise AssertionError("global typed serde fallback must not be used")
repo = _Repo()
saver = RepositoryCheckpointSaver(SimpleNamespace(), repository=repo)
saver.serde = ExplodingSerde()
config = {"configurable": {"thread_id": "bad-checkpoint-thread"}}
with pytest.raises(TypeError, match="Checkpoint contém valor não serializável"):
asyncio.run(saver.aput(
config,
{"id": "cp-bad", "channel_values": {"bad": RuntimeLike()}},
{},
{},
))
def test_json_pending_writes_remain_plain_json() -> None:
repo = _Repo()
saver = RepositoryCheckpointSaver(SimpleNamespace(), repository=repo)
config = {"configurable": {"thread_id": "plain-thread"}}
asyncio.run(saver.aput(config, {"id": "cp-plain", "channel_values": {}}, {}, {}))
asyncio.run(saver.aput_writes(config, [("result", {"ok": True})], "task-plain"))
persisted = repo.saved[1]
assert persisted["pending_writes"][0]["value"] == {"ok": True}

View File

@@ -0,0 +1,117 @@
from types import SimpleNamespace
import pytest
from agent_framework.rag.rag_service import RagService
def _settings(**overrides):
data = dict(
RAG_PROVIDER="kbdb", RAG_TOP_K=5,
KBDB_DB_USER="kb_user", KBDB_DB_PASSWORD="kb_pwd", KBDB_DB_DSN="kb_tp",
KBDB_DB_WALLET_LOCATION=None, KBDB_DB_WALLET_PASSWORD=None,
ADB_USER=None, ADB_PASSWORD=None, ADB_DSN=None,
ADB_WALLET_LOCATION=None, ADB_WALLET_PASSWORD=None,
KBDB_SEARCH_TYPE="hybrid", KBDB_NODE_EXPANSION=True,
KBDB_NODE_MAX_RELATED=8, KBDB_GRAPH_CROSS_REF=False,
KBDB_MAX_CROSS_REF_HOPS=1, KBDB_DOCUMENT_TYPE="customer_safe",
KBDB_METADATA_JSON=None, KBDB_MIN_SCORE=None,
)
data.update(overrides)
return SimpleNamespace(**data)
@pytest.mark.asyncio
async def test_kbdb_provider_adapts_serving_envelope(monkeypatch):
service = RagService(_settings())
def fake_search(query, k):
return {
"search_type": "hybrid", "confidence": "high", "low_confidence": False,
"top_score": 78.4, "warnings": [],
"seeds": [{"unit_id": 11, "rank": 1, "score": 78.4}],
"units": [
{"unit_id": 10, "content": "passo anterior", "provenance": "parent"},
{"unit_id": 11, "content": "resposta principal", "provenance": "seed"},
],
"documents": [{"document_id": 7, "title": "Politica"}],
}
monkeypatch.setattr(service._kbdb, "_search_sync", fake_search)
result = await service.retrieve("qual a regra?", namespace="billing_agent")
assert [d.id for d in result.documents] == ["10", "11"]
assert result.documents[1].score == 78.4
assert result.metadata["provider"] == "kbdb"
assert result.metadata["confidence"] == "high"
assert "resposta principal" in result.as_prompt_context()
@pytest.mark.asyncio
async def test_kbdb_provider_is_serving_only():
service = RagService(_settings())
with pytest.raises(RuntimeError, match="serving-only"):
await service.add_documents(["texto"])
def test_kbdb_connection_uses_same_wallet_semantics_without_adb_fallback(monkeypatch):
import sys
from agent_framework.rag.kbdb_service import KbdbRagService
captured = {}
class Defaults:
fetch_lobs = True
class Connection:
def close(self):
captured["closed"] = True
class FakeOracleDb:
defaults = Defaults()
@staticmethod
def connect(**kwargs):
captured.update(kwargs)
return Connection()
monkeypatch.setitem(sys.modules, "oracledb", FakeOracleDb)
settings = _settings(
KBDB_DB_USER="kb_user",
KBDB_DB_PASSWORD="kb_pwd",
KBDB_DB_DSN="kb_tp",
KBDB_DB_WALLET_LOCATION="/wallet/kb",
KBDB_DB_WALLET_PASSWORD="wallet_pwd",
ADB_USER="framework_user",
ADB_PASSWORD="framework_pwd",
ADB_DSN="framework_high",
ADB_WALLET_LOCATION="/wallet/framework",
ADB_WALLET_PASSWORD="framework_wallet_pwd",
)
service = KbdbRagService(settings)
with service._connect():
pass
assert captured["user"] == "kb_user"
assert captured["password"] == "kb_pwd"
assert captured["dsn"] == "kb_tp"
assert captured["config_dir"] == "/wallet/kb"
assert captured["wallet_location"] == "/wallet/kb"
assert captured["wallet_password"] == "wallet_pwd"
assert captured["closed"] is True
def test_kbdb_does_not_fallback_to_framework_adb_credentials():
from agent_framework.rag.kbdb_service import KbdbRagService
settings = _settings(
KBDB_DB_USER=None,
KBDB_DB_PASSWORD=None,
KBDB_DB_DSN=None,
ADB_USER="framework_user",
ADB_PASSWORD="framework_pwd",
ADB_DSN="framework_high",
)
with pytest.raises(RuntimeError, match="KBDB_DB_USER"):
KbdbRagService(settings)

View File

@@ -0,0 +1,100 @@
from types import SimpleNamespace
import pytest
from agent_framework.rag.rag_service import RagResult
from agent_framework.rag.vector_store import VectorDocument
from agent_framework.runtime.agent_runtime import AgentRuntimeMixin
class DummyRag:
def __init__(self):
self.calls = []
async def retrieve(self, query, *, namespace="default", graph_node=None, rewrite=False, k=None):
self.calls.append((query, namespace, graph_node, rewrite))
return RagResult(
query=query,
documents=[VectorDocument(id="kb-1", content="Tarifação documentada", metadata={}, score=0.9)],
graph_neighbors=[],
latency_ms=3,
metadata={"provider": "kbdb", "confidence": "high", "low_confidence": False},
)
class Runtime(AgentRuntimeMixin):
pass
def _runtime(**settings):
rt = Runtime()
base = dict(
RAG_PROVIDER="kbdb",
SKIP_RAG_WHEN_MCP_SUFFICIENT=True,
ENABLE_RAG_QUERY_REWRITE=False,
ENABLE_RAG_CONTEXT_COMPRESSION=False,
RAG_GROUNDED_ONLY=False,
KBDB_GROUNDED_ONLY=True,
LONG_TERM_MEMORY_INJECT_CONTEXT=False,
)
base.update(settings)
rt.settings = SimpleNamespace(**base)
rt.rag_service = DummyRag()
rt.guardrail_pipeline = None
return rt
@pytest.mark.asyncio
async def test_successful_mcp_does_not_skip_rag_without_explicit_sufficiency():
rt = _runtime()
state = {
"agent_id": "telecom_contas",
"user_text": "Como funciona a tarifação do plano Infinity Pós?",
"sanitized_input": "Como funciona a tarifação do plano Infinity Pós?",
"mcp_results": [{"ok": True, "tool_name": "qualquer_tool", "result": {"plano": "Controle 50GB"}}],
}
context, metadata = await rt._retrieve_rag_context(state)
assert rt.rag_service.calls
assert "Tarifação documentada" in context
assert metadata["provider"] == "kbdb"
assert metadata["status"] == "executed"
assert metadata["document_count"] == 1
@pytest.mark.asyncio
async def test_rag_skips_only_when_mcp_explicitly_declares_sufficiency():
rt = _runtime()
state = {
"agent_id": "telecom_contas",
"user_text": "qual é meu plano?",
"sanitized_input": "qual é meu plano?",
"mcp_results": [{"ok": True, "result": {"plano": "Controle 50GB", "rag_sufficient": True}}],
}
context, metadata = await rt._retrieve_rag_context(state)
assert context == ""
assert not rt.rag_service.calls
assert metadata["reason"] == "mcp_explicitly_sufficient"
def test_kbdb_build_messages_injects_grounding_policy():
rt = _runtime()
state = {
"user_text": "Como funciona?",
"sanitized_input": "Como funciona?",
"context": {},
"business_context": {},
}
messages = rt.build_messages(
state,
system_prompt="system",
mcp_results=[{"ok": True, "result": {"plano": "Controle"}}],
rag_context="",
rag_metadata={"provider": "kbdb", "enabled": True, "status": "empty", "document_count": 0},
)
user = next(m["content"] for m in messages if m["role"] == "user")
assert "Política de grounding obrigatória" in user
assert "Não complete lacunas usando conhecimento paramétrico" in user