mirror of
https://github.com/hoshikawa2/agent_platform_oci.git
synced 2026-09-07 18:23:46 +00:00
new feature: Integration with kbdb Autonomous
This commit is contained in:
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
160
tests/unit/test_langgraph_checkpoint_interrupt_controlled.py
Normal file
160
tests/unit/test_langgraph_checkpoint_interrupt_controlled.py
Normal 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()},
|
||||
{},
|
||||
))
|
||||
@@ -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}
|
||||
|
||||
117
tests/unit/test_rag_kbdb_provider.py
Normal file
117
tests/unit/test_rag_kbdb_provider.py
Normal 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)
|
||||
100
tests/unit/test_rag_runtime_grounding.py
Normal file
100
tests/unit/test_rag_runtime_grounding.py
Normal 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
|
||||
Reference in New Issue
Block a user