topoteretes--cognee
c889a57b6b
Test Suites / Build CI Environment (push) Has been cancelled
Test Suites / Basic Tests (push) Has been cancelled
Test Suites / End-to-End Tests (push) Has been cancelled
Test Suites / CLI Tests (push) Has been cancelled
Test Suites / Slow End-to-End Tests (push) Has been cancelled
Test Suites / Graph Database Tests (push) Has been cancelled
Test Suites / Vector DB Tests (push) Has been cancelled
Test Suites / Temporal Graph Test (push) Has been cancelled
Test Suites / Search Test on Different DBs (push) Has been cancelled
Test Suites / Example Tests (push) Has been cancelled
Test Suites / Notebook Tests (push) Has been cancelled
Test Suites / OS and Python Tests Ubuntu (push) Has been cancelled
Test Suites / OS and Python Tests Extended (push) Has been cancelled
Test Suites / LLM Test Suite (push) Has been cancelled
Test Suites / S3 File Storage Test (push) Has been cancelled
Test Suites / Run Integration Tests (push) Has been cancelled
Test Suites / MCP Tests (push) Has been cancelled
Test Suites / Docker Compose Test (push) Has been cancelled
Test Suites / Docker CI test (push) Has been cancelled
Test Suites / Relational DB Migration Tests (push) Has been cancelled
Test Suites / Distributed Cognee Test (push) Has been cancelled
Test Suites / DB Examples Tests (push) Has been cancelled
Test Suites / Test Completion Status (push) Has been cancelled
Test Suites / Claude Code Review (push) Has been cancelled
Test Suites / basic checks (push) Has been cancelled
build | Build and Push Cognee MCP Docker Image to dockerhub / docker-build-and-push (push) Has been cancelled
Scorecard supply-chain security / Scorecard analysis (push) Has been cancelled
build | Build and Push Docker Image to dockerhub / docker-build-and-push (push) Has been cancelled
Weighted Edges Tests / Test Weighted Edges Core Functionality (3.11) (push) Has been cancelled
Weighted Edges Tests / Test Weighted Edges Core Functionality (3.12) (push) Has been cancelled
Weighted Edges Tests / Test Weighted Edges with Different Graph Databases (kuzu, kuzu) (push) Has been cancelled
Weighted Edges Tests / Test Weighted Edges with Different Graph Databases (neo4j, neo4j) (push) Has been cancelled
Weighted Edges Tests / Test Weighted Edges Examples (push) Has been cancelled
Weighted Edges Tests / Code Quality for Weighted Edges (push) Has been cancelled
368 行
15 KiB
Python
368 行
15 KiB
Python
"""Remote HTTP client that proxies V2 operations to a Cognee Cloud instance."""
|
|
|
|
import io
|
|
from pathlib import Path
|
|
from typing import Any, Optional
|
|
from uuid import UUID
|
|
|
|
import aiohttp
|
|
|
|
from cognee.shared.logging_utils import get_logger
|
|
|
|
logger = get_logger("serve.cloud_client")
|
|
|
|
|
|
class CloudClient:
|
|
"""Async HTTP client for a remote Cognee Cloud tenant instance.
|
|
|
|
All requests use ``X-Api-Key`` for authentication, matching the
|
|
SaaS backend's API key auth backend.
|
|
"""
|
|
|
|
def __init__(self, service_url: str, api_key: str):
|
|
self.service_url = service_url.rstrip("/")
|
|
self.api_key = api_key
|
|
self._session: Optional[aiohttp.ClientSession] = None
|
|
|
|
# Default for ordinary API calls: aiohttp's standard 5-minute total,
|
|
# with connect failures surfacing quickly.
|
|
DEFAULT_TIMEOUT = aiohttp.ClientTimeout(total=300, sock_connect=30)
|
|
# Archive uploads (cognee.push) plus the synchronous server-side import
|
|
# can legitimately exceed any fixed total; per-read inactivity stays
|
|
# bounded instead. Applied per-request, only to archive uploads.
|
|
UPLOAD_TIMEOUT = aiohttp.ClientTimeout(total=None, sock_connect=30, sock_read=300)
|
|
|
|
async def _get_session(self) -> aiohttp.ClientSession:
|
|
if self._session is None or self._session.closed:
|
|
self._session = aiohttp.ClientSession(
|
|
headers={"X-Api-Key": self.api_key},
|
|
timeout=self.DEFAULT_TIMEOUT,
|
|
)
|
|
return self._session
|
|
|
|
async def close(self) -> None:
|
|
if self._session and not self._session.closed:
|
|
await self._session.close()
|
|
self._session = None
|
|
|
|
async def _health_check(self) -> bool:
|
|
"""Verify the remote instance is reachable."""
|
|
try:
|
|
session = await self._get_session()
|
|
async with session.get(f"{self.service_url}/health") as resp:
|
|
return resp.status == 200
|
|
except Exception:
|
|
return False
|
|
|
|
# ----- V2 Operations -----
|
|
|
|
async def remember(self, data: Any, dataset_name: str = "main_dataset", **kwargs) -> dict:
|
|
"""POST /api/v1/remember — ingest data and build knowledge graph."""
|
|
session = await self._get_session()
|
|
|
|
form = aiohttp.FormData()
|
|
form.add_field("datasetName", dataset_name)
|
|
|
|
if kwargs.get("session_id"):
|
|
form.add_field("session_id", kwargs["session_id"])
|
|
if kwargs.get("run_in_background"):
|
|
form.add_field("run_in_background", "true")
|
|
if kwargs.get("custom_prompt"):
|
|
form.add_field("custom_prompt", kwargs["custom_prompt"])
|
|
if kwargs.get("chunk_size") is not None:
|
|
form.add_field("chunk_size", str(kwargs["chunk_size"]))
|
|
if kwargs.get("chunks_per_batch") is not None:
|
|
form.add_field("chunks_per_batch", str(kwargs["chunks_per_batch"]))
|
|
content_type_kw = kwargs.get("content_type")
|
|
if content_type_kw is not None:
|
|
form.add_field("content_type", str(content_type_kw))
|
|
if kwargs.get("import_mode") is not None:
|
|
form.add_field("import_mode", str(kwargs["import_mode"]))
|
|
|
|
# Skills are local SKILL.md files. The server's add_skills() reads
|
|
# paths from its own filesystem — sending the path string verbatim
|
|
# would have the server look for that path on the POD, not the
|
|
# caller. For content_type="skills", read each SKILL.md and upload
|
|
# its bytes so the server can write them to a tempdir.
|
|
if content_type_kw == "skills" and isinstance(data, (str, Path)):
|
|
source = Path(data).expanduser()
|
|
if source.is_file():
|
|
skill_files = [source] if source.name == "SKILL.md" else []
|
|
elif source.is_dir():
|
|
skill_files = sorted(source.rglob("SKILL.md"))
|
|
else:
|
|
raise FileNotFoundError(f"Skills source not found: {data}")
|
|
if not skill_files:
|
|
raise ValueError(f"No SKILL.md files under {data}")
|
|
base = source if source.is_dir() else source.parent
|
|
for skill_path in skill_files:
|
|
# Preserve relative structure so the server can reconstruct
|
|
# the SKILL.md layout when writing to its tempdir.
|
|
rel = skill_path.relative_to(base).as_posix()
|
|
form.add_field("data", skill_path.open("rb"), filename=rel)
|
|
# Handle data — string or file-like objects
|
|
elif isinstance(data, str):
|
|
form.add_field(
|
|
"data",
|
|
io.BytesIO(data.encode("utf-8")),
|
|
filename="data.txt",
|
|
content_type="text/plain",
|
|
)
|
|
elif isinstance(data, list):
|
|
for item in data:
|
|
if isinstance(item, str):
|
|
form.add_field(
|
|
"data",
|
|
io.BytesIO(item.encode("utf-8")),
|
|
filename="data.txt",
|
|
content_type="text/plain",
|
|
)
|
|
elif hasattr(item, "read"):
|
|
name = getattr(item, "name", "upload")
|
|
form.add_field("data", item, filename=name)
|
|
elif hasattr(data, "read"):
|
|
name = getattr(data, "name", "upload")
|
|
form.add_field("data", data, filename=name)
|
|
|
|
timeout = (
|
|
self.UPLOAD_TIMEOUT
|
|
if kwargs.get("content_type") == "cogx-archive"
|
|
else self.DEFAULT_TIMEOUT
|
|
)
|
|
async with session.post(
|
|
f"{self.service_url}/api/v1/remember", data=form, timeout=timeout
|
|
) as resp:
|
|
if resp.status >= 400:
|
|
body = await resp.text()
|
|
raise RuntimeError(f"Remote remember failed ({resp.status}): {body}")
|
|
return await resp.json()
|
|
|
|
async def remember_entry(
|
|
self,
|
|
entry,
|
|
dataset_name: str = "main_dataset",
|
|
session_id: Optional[str] = None,
|
|
skill_improvement: Optional[dict] = None,
|
|
) -> dict:
|
|
"""POST /api/v1/remember/entry — store a typed MemoryEntry.
|
|
|
|
``entry`` is a pydantic MemoryEntry.
|
|
"""
|
|
session = await self._get_session()
|
|
|
|
# Pydantic v2: model_dump preserves the discriminator field.
|
|
entry_dump = entry.model_dump(mode="json")
|
|
|
|
payload = {
|
|
"entry": entry_dump,
|
|
"dataset_name": dataset_name,
|
|
"session_id": session_id,
|
|
"skill_improvement": skill_improvement,
|
|
}
|
|
|
|
async with session.post(
|
|
f"{self.service_url}/api/v1/remember/entry",
|
|
json=payload,
|
|
) as resp:
|
|
if resp.status >= 400:
|
|
body = await resp.text()
|
|
raise RuntimeError(f"Remote remember_entry failed ({resp.status}): {body}")
|
|
return await resp.json()
|
|
|
|
async def recall(self, query_text: str, query_type: Optional[str] = None, **kwargs) -> list:
|
|
"""POST /api/v1/recall — query the knowledge graph and/or session cache."""
|
|
session = await self._get_session()
|
|
|
|
payload: dict = {"query": query_text}
|
|
if query_type:
|
|
payload["search_type"] = query_type if isinstance(query_type, str) else query_type.value
|
|
if kwargs.get("dataset_ids"):
|
|
payload["dataset_ids"] = [str(dataset_id) for dataset_id in kwargs["dataset_ids"]]
|
|
elif kwargs.get("datasets"):
|
|
payload["datasets"] = kwargs["datasets"]
|
|
if kwargs.get("top_k"):
|
|
payload["top_k"] = kwargs["top_k"]
|
|
if kwargs.get("system_prompt"):
|
|
payload["system_prompt"] = kwargs["system_prompt"]
|
|
if kwargs.get("node_name"):
|
|
payload["node_name"] = kwargs["node_name"]
|
|
if kwargs.get("only_context"):
|
|
payload["only_context"] = kwargs["only_context"]
|
|
if kwargs.get("verbose"):
|
|
payload["verbose"] = kwargs["verbose"]
|
|
if kwargs.get("session_id"):
|
|
payload["session_id"] = kwargs["session_id"]
|
|
if kwargs.get("scope") is not None:
|
|
payload["scope"] = kwargs["scope"]
|
|
if kwargs.get("context_profile") is not None:
|
|
payload["context_profile"] = kwargs["context_profile"]
|
|
if kwargs.get("include_references") is not None:
|
|
payload["include_references"] = kwargs["include_references"]
|
|
|
|
async with session.post(
|
|
f"{self.service_url}/api/v1/recall",
|
|
json=payload,
|
|
) as resp:
|
|
if resp.status >= 400:
|
|
body = await resp.text()
|
|
raise RuntimeError(f"Remote recall failed ({resp.status}): {body}")
|
|
return await resp.json()
|
|
|
|
async def improve(self, dataset: Any = "main_dataset", **kwargs) -> dict:
|
|
"""POST /api/v1/improve — enrich the knowledge graph."""
|
|
session = await self._get_session()
|
|
|
|
payload = {}
|
|
if isinstance(dataset, UUID):
|
|
payload["dataset_id"] = str(dataset)
|
|
else:
|
|
payload["dataset_name"] = str(dataset)
|
|
if kwargs.get("run_in_background"):
|
|
payload["run_in_background"] = True
|
|
if kwargs.get("node_name"):
|
|
payload["node_name"] = kwargs["node_name"]
|
|
|
|
async with session.post(
|
|
f"{self.service_url}/api/v1/improve",
|
|
json=payload,
|
|
) as resp:
|
|
if resp.status >= 400:
|
|
body = await resp.text()
|
|
raise RuntimeError(f"Remote improve failed ({resp.status}): {body}")
|
|
return await resp.json()
|
|
|
|
# ----- V1 Operations (add / cognify / search) -----
|
|
|
|
async def add(self, data: Any, dataset_name: str = "main_dataset", **kwargs) -> dict:
|
|
"""POST /api/v1/add — ingest data into a dataset."""
|
|
session = await self._get_session()
|
|
|
|
form = aiohttp.FormData()
|
|
form.add_field("datasetName", dataset_name)
|
|
|
|
if isinstance(data, str):
|
|
form.add_field(
|
|
"data",
|
|
io.BytesIO(data.encode("utf-8")),
|
|
filename="data.txt",
|
|
content_type="text/plain",
|
|
)
|
|
elif isinstance(data, list):
|
|
for item in data:
|
|
if isinstance(item, str):
|
|
form.add_field(
|
|
"data",
|
|
io.BytesIO(item.encode("utf-8")),
|
|
filename="data.txt",
|
|
content_type="text/plain",
|
|
)
|
|
elif hasattr(item, "read"):
|
|
name = getattr(item, "name", "upload")
|
|
form.add_field("data", item, filename=name)
|
|
elif hasattr(data, "read"):
|
|
name = getattr(data, "name", "upload")
|
|
form.add_field("data", data, filename=name)
|
|
|
|
async with session.post(f"{self.service_url}/api/v1/add", data=form) as resp:
|
|
if resp.status >= 400:
|
|
body = await resp.text()
|
|
raise RuntimeError(f"Remote add failed ({resp.status}): {body}")
|
|
return await resp.json()
|
|
|
|
async def cognify(self, datasets: Any = None, **kwargs) -> dict:
|
|
"""POST /api/v1/cognify — build the knowledge graph."""
|
|
session = await self._get_session()
|
|
|
|
payload: dict = {}
|
|
if datasets:
|
|
payload["datasets"] = (
|
|
[str(d) for d in datasets] if isinstance(datasets, list) else [str(datasets)]
|
|
)
|
|
if kwargs.get("run_in_background"):
|
|
payload["run_in_background"] = True
|
|
if kwargs.get("custom_prompt"):
|
|
payload["custom_prompt"] = kwargs["custom_prompt"]
|
|
if kwargs.get("chunk_size") is not None:
|
|
payload["chunk_size"] = kwargs["chunk_size"]
|
|
if kwargs.get("chunks_per_batch") is not None:
|
|
payload["chunks_per_batch"] = kwargs["chunks_per_batch"]
|
|
|
|
async with session.post(
|
|
f"{self.service_url}/api/v1/cognify",
|
|
json=payload,
|
|
) as resp:
|
|
if resp.status >= 400:
|
|
body = await resp.text()
|
|
raise RuntimeError(f"Remote cognify failed ({resp.status}): {body}")
|
|
return await resp.json()
|
|
|
|
async def search(self, query: str, **kwargs) -> list:
|
|
"""POST /api/v1/search — query the knowledge graph."""
|
|
session = await self._get_session()
|
|
|
|
payload: dict = {"query": query}
|
|
if kwargs.get("search_type"):
|
|
st = kwargs["search_type"]
|
|
payload["searchType"] = st if isinstance(st, str) else st.value
|
|
if kwargs.get("datasets"):
|
|
payload["datasets"] = kwargs["datasets"]
|
|
if kwargs.get("dataset_ids"):
|
|
dataset_ids = kwargs["dataset_ids"]
|
|
if isinstance(dataset_ids, UUID):
|
|
dataset_ids = [dataset_ids]
|
|
payload["datasetIds"] = [str(dataset_id) for dataset_id in dataset_ids]
|
|
if kwargs.get("top_k") is not None:
|
|
payload["topK"] = kwargs["top_k"]
|
|
if kwargs.get("system_prompt"):
|
|
payload["systemPrompt"] = kwargs["system_prompt"]
|
|
if kwargs.get("node_name"):
|
|
payload["nodeName"] = kwargs["node_name"]
|
|
if kwargs.get("only_context") is not None:
|
|
payload["onlyContext"] = kwargs["only_context"]
|
|
if kwargs.get("verbose") is not None:
|
|
payload["verbose"] = kwargs["verbose"]
|
|
if kwargs.get("skills") is not None:
|
|
payload["skills"] = [
|
|
skill.name if hasattr(skill, "name") else str(skill) for skill in kwargs["skills"]
|
|
]
|
|
if kwargs.get("tools") is not None:
|
|
payload["tools"] = kwargs["tools"]
|
|
if kwargs.get("max_iter") is not None:
|
|
payload["maxIter"] = kwargs["max_iter"]
|
|
if kwargs.get("include_references") is not None:
|
|
payload["includeReferences"] = kwargs["include_references"]
|
|
|
|
async with session.post(
|
|
f"{self.service_url}/api/v1/search",
|
|
json=payload,
|
|
) as resp:
|
|
if resp.status >= 400:
|
|
body = await resp.text()
|
|
raise RuntimeError(f"Remote search failed ({resp.status}): {body}")
|
|
return await resp.json()
|
|
|
|
async def forget(self, **kwargs) -> dict:
|
|
"""POST /api/v1/forget — delete data from the knowledge graph."""
|
|
session = await self._get_session()
|
|
|
|
payload = {}
|
|
if kwargs.get("everything"):
|
|
payload["everything"] = True
|
|
if kwargs.get("dataset"):
|
|
payload["dataset"] = str(kwargs["dataset"])
|
|
if kwargs.get("dataset_id"):
|
|
payload["dataset_id"] = str(kwargs["dataset_id"])
|
|
if kwargs.get("data_id"):
|
|
payload["data_id"] = str(kwargs["data_id"])
|
|
if kwargs.get("memory_only") is not None:
|
|
payload["memory_only"] = bool(kwargs["memory_only"])
|
|
|
|
async with session.post(
|
|
f"{self.service_url}/api/v1/forget",
|
|
json=payload,
|
|
) as resp:
|
|
if resp.status >= 400:
|
|
body = await resp.text()
|
|
raise RuntimeError(f"Remote forget failed ({resp.status}): {body}")
|
|
return await resp.json()
|