mirror of
https://github.com/trustgraph-ai/trustgraph.git
synced 2026-07-22 19:51:02 +02:00
CLI tools for tg-invoke-graph-embeddings, tg-invoke-document-embeddings,
and tg-invoke-embeddings. Just useful for diagnostics.
This commit is contained in:
parent
23cc4dfdd1
commit
cd91262e9d
8 changed files with 500 additions and 18 deletions
|
|
@ -612,8 +612,12 @@ class AsyncFlowInstance:
|
||||||
print(f"{entity['name']}: {entity['score']}")
|
print(f"{entity['name']}: {entity['score']}")
|
||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
|
# First convert text to embeddings vectors
|
||||||
|
emb_result = await self.embeddings(text=text)
|
||||||
|
vectors = emb_result.get("vectors", [])
|
||||||
|
|
||||||
request_data = {
|
request_data = {
|
||||||
"text": text,
|
"vectors": vectors,
|
||||||
"user": user,
|
"user": user,
|
||||||
"collection": collection,
|
"collection": collection,
|
||||||
"limit": limit
|
"limit": limit
|
||||||
|
|
|
||||||
|
|
@ -282,8 +282,12 @@ class AsyncSocketFlowInstance:
|
||||||
|
|
||||||
async def graph_embeddings_query(self, text: str, user: str, collection: str, limit: int = 10, **kwargs):
|
async def graph_embeddings_query(self, text: str, user: str, collection: str, limit: int = 10, **kwargs):
|
||||||
"""Query graph embeddings for semantic search"""
|
"""Query graph embeddings for semantic search"""
|
||||||
|
# First convert text to embeddings vectors
|
||||||
|
emb_result = await self.embeddings(text=text)
|
||||||
|
vectors = emb_result.get("vectors", [])
|
||||||
|
|
||||||
request = {
|
request = {
|
||||||
"text": text,
|
"vectors": vectors,
|
||||||
"user": user,
|
"user": user,
|
||||||
"collection": collection,
|
"collection": collection,
|
||||||
"limit": limit
|
"limit": limit
|
||||||
|
|
|
||||||
|
|
@ -15,6 +15,15 @@ from . types import Triple
|
||||||
from . exceptions import ProtocolException
|
from . exceptions import ProtocolException
|
||||||
|
|
||||||
|
|
||||||
|
def _string_to_term(value: str) -> Dict[str, Any]:
|
||||||
|
"""Convert a string value to Term format for the gateway."""
|
||||||
|
# Treat URIs as IRI type, otherwise as literal
|
||||||
|
if value.startswith("http://") or value.startswith("https://") or "://" in value:
|
||||||
|
return {"t": "i", "i": value}
|
||||||
|
else:
|
||||||
|
return {"t": "l", "v": value}
|
||||||
|
|
||||||
|
|
||||||
class BulkClient:
|
class BulkClient:
|
||||||
"""
|
"""
|
||||||
Synchronous bulk operations client for import/export.
|
Synchronous bulk operations client for import/export.
|
||||||
|
|
@ -62,7 +71,12 @@ class BulkClient:
|
||||||
|
|
||||||
return loop.run_until_complete(coro)
|
return loop.run_until_complete(coro)
|
||||||
|
|
||||||
def import_triples(self, flow: str, triples: Iterator[Triple], **kwargs: Any) -> None:
|
def import_triples(
|
||||||
|
self, flow: str, triples: Iterator[Triple],
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
batch_size: int = 100,
|
||||||
|
**kwargs: Any
|
||||||
|
) -> None:
|
||||||
"""
|
"""
|
||||||
Bulk import RDF triples into a flow.
|
Bulk import RDF triples into a flow.
|
||||||
|
|
||||||
|
|
@ -71,6 +85,8 @@ class BulkClient:
|
||||||
Args:
|
Args:
|
||||||
flow: Flow identifier
|
flow: Flow identifier
|
||||||
triples: Iterator yielding Triple objects
|
triples: Iterator yielding Triple objects
|
||||||
|
metadata: Metadata dict with id, metadata, user, collection
|
||||||
|
batch_size: Number of triples per batch (default 100)
|
||||||
**kwargs: Additional parameters (reserved for future use)
|
**kwargs: Additional parameters (reserved for future use)
|
||||||
|
|
||||||
Example:
|
Example:
|
||||||
|
|
@ -86,23 +102,47 @@ class BulkClient:
|
||||||
# ... more triples
|
# ... more triples
|
||||||
|
|
||||||
# Import triples
|
# Import triples
|
||||||
bulk.import_triples(flow="default", triples=triple_generator())
|
bulk.import_triples(
|
||||||
|
flow="default",
|
||||||
|
triples=triple_generator(),
|
||||||
|
metadata={"id": "doc1", "metadata": [], "user": "user1", "collection": "default"}
|
||||||
|
)
|
||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
self._run_async(self._import_triples_async(flow, triples))
|
self._run_async(self._import_triples_async(flow, triples, metadata, batch_size))
|
||||||
|
|
||||||
async def _import_triples_async(self, flow: str, triples: Iterator[Triple]) -> None:
|
async def _import_triples_async(
|
||||||
|
self, flow: str, triples: Iterator[Triple],
|
||||||
|
metadata: Optional[Dict[str, Any]], batch_size: int
|
||||||
|
) -> None:
|
||||||
"""Async implementation of triple import"""
|
"""Async implementation of triple import"""
|
||||||
ws_url = f"{self.url}/api/v1/flow/{flow}/import/triples"
|
ws_url = f"{self.url}/api/v1/flow/{flow}/import/triples"
|
||||||
if self.token:
|
if self.token:
|
||||||
ws_url = f"{ws_url}?token={self.token}"
|
ws_url = f"{ws_url}?token={self.token}"
|
||||||
|
|
||||||
|
if metadata is None:
|
||||||
|
metadata = {"id": "", "metadata": [], "user": "trustgraph", "collection": "default"}
|
||||||
|
|
||||||
async with websockets.connect(ws_url, ping_interval=20, ping_timeout=self.timeout) as websocket:
|
async with websockets.connect(ws_url, ping_interval=20, ping_timeout=self.timeout) as websocket:
|
||||||
|
batch = []
|
||||||
for triple in triples:
|
for triple in triples:
|
||||||
|
batch.append({
|
||||||
|
"s": _string_to_term(triple.s),
|
||||||
|
"p": _string_to_term(triple.p),
|
||||||
|
"o": _string_to_term(triple.o)
|
||||||
|
})
|
||||||
|
if len(batch) >= batch_size:
|
||||||
|
message = {
|
||||||
|
"metadata": metadata,
|
||||||
|
"triples": batch
|
||||||
|
}
|
||||||
|
await websocket.send(json.dumps(message))
|
||||||
|
batch = []
|
||||||
|
# Send remaining items
|
||||||
|
if batch:
|
||||||
message = {
|
message = {
|
||||||
"s": triple.s,
|
"metadata": metadata,
|
||||||
"p": triple.p,
|
"triples": batch
|
||||||
"o": triple.o
|
|
||||||
}
|
}
|
||||||
await websocket.send(json.dumps(message))
|
await websocket.send(json.dumps(message))
|
||||||
|
|
||||||
|
|
@ -362,7 +402,12 @@ class BulkClient:
|
||||||
async for raw_message in websocket:
|
async for raw_message in websocket:
|
||||||
yield json.loads(raw_message)
|
yield json.loads(raw_message)
|
||||||
|
|
||||||
def import_entity_contexts(self, flow: str, contexts: Iterator[Dict[str, Any]], **kwargs: Any) -> None:
|
def import_entity_contexts(
|
||||||
|
self, flow: str, contexts: Iterator[Dict[str, Any]],
|
||||||
|
metadata: Optional[Dict[str, Any]] = None,
|
||||||
|
batch_size: int = 100,
|
||||||
|
**kwargs: Any
|
||||||
|
) -> None:
|
||||||
"""
|
"""
|
||||||
Bulk import entity contexts into a flow.
|
Bulk import entity contexts into a flow.
|
||||||
|
|
||||||
|
|
@ -373,6 +418,8 @@ class BulkClient:
|
||||||
Args:
|
Args:
|
||||||
flow: Flow identifier
|
flow: Flow identifier
|
||||||
contexts: Iterator yielding context dictionaries
|
contexts: Iterator yielding context dictionaries
|
||||||
|
metadata: Metadata dict with id, metadata, user, collection
|
||||||
|
batch_size: Number of contexts per batch (default 100)
|
||||||
**kwargs: Additional parameters (reserved for future use)
|
**kwargs: Additional parameters (reserved for future use)
|
||||||
|
|
||||||
Example:
|
Example:
|
||||||
|
|
@ -381,27 +428,49 @@ class BulkClient:
|
||||||
|
|
||||||
# Generate entity contexts to import
|
# Generate entity contexts to import
|
||||||
def context_generator():
|
def context_generator():
|
||||||
yield {"entity": "entity1", "context": "Description of entity1..."}
|
yield {"entity": {"v": "entity1", "e": True}, "context": "Description..."}
|
||||||
yield {"entity": "entity2", "context": "Description of entity2..."}
|
yield {"entity": {"v": "entity2", "e": True}, "context": "Description..."}
|
||||||
# ... more contexts
|
# ... more contexts
|
||||||
|
|
||||||
bulk.import_entity_contexts(
|
bulk.import_entity_contexts(
|
||||||
flow="default",
|
flow="default",
|
||||||
contexts=context_generator()
|
contexts=context_generator(),
|
||||||
|
metadata={"id": "doc1", "metadata": [], "user": "user1", "collection": "default"}
|
||||||
)
|
)
|
||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
self._run_async(self._import_entity_contexts_async(flow, contexts))
|
self._run_async(self._import_entity_contexts_async(flow, contexts, metadata, batch_size))
|
||||||
|
|
||||||
async def _import_entity_contexts_async(self, flow: str, contexts: Iterator[Dict[str, Any]]) -> None:
|
async def _import_entity_contexts_async(
|
||||||
|
self, flow: str, contexts: Iterator[Dict[str, Any]],
|
||||||
|
metadata: Optional[Dict[str, Any]], batch_size: int
|
||||||
|
) -> None:
|
||||||
"""Async implementation of entity contexts import"""
|
"""Async implementation of entity contexts import"""
|
||||||
ws_url = f"{self.url}/api/v1/flow/{flow}/import/entity-contexts"
|
ws_url = f"{self.url}/api/v1/flow/{flow}/import/entity-contexts"
|
||||||
if self.token:
|
if self.token:
|
||||||
ws_url = f"{ws_url}?token={self.token}"
|
ws_url = f"{ws_url}?token={self.token}"
|
||||||
|
|
||||||
|
if metadata is None:
|
||||||
|
metadata = {"id": "", "metadata": [], "user": "trustgraph", "collection": "default"}
|
||||||
|
|
||||||
async with websockets.connect(ws_url, ping_interval=20, ping_timeout=self.timeout) as websocket:
|
async with websockets.connect(ws_url, ping_interval=20, ping_timeout=self.timeout) as websocket:
|
||||||
|
batch = []
|
||||||
for context in contexts:
|
for context in contexts:
|
||||||
await websocket.send(json.dumps(context))
|
batch.append(context)
|
||||||
|
if len(batch) >= batch_size:
|
||||||
|
message = {
|
||||||
|
"metadata": metadata,
|
||||||
|
"entities": batch
|
||||||
|
}
|
||||||
|
await websocket.send(json.dumps(message))
|
||||||
|
batch = []
|
||||||
|
# Send remaining items
|
||||||
|
if batch:
|
||||||
|
message = {
|
||||||
|
"metadata": metadata,
|
||||||
|
"entities": batch
|
||||||
|
}
|
||||||
|
await websocket.send(json.dumps(message))
|
||||||
|
|
||||||
def export_entity_contexts(self, flow: str, **kwargs: Any) -> Iterator[Dict[str, Any]]:
|
def export_entity_contexts(self, flow: str, **kwargs: Any) -> Iterator[Dict[str, Any]]:
|
||||||
"""
|
"""
|
||||||
|
|
|
||||||
|
|
@ -584,9 +584,13 @@ class FlowInstance:
|
||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
# First convert text to embeddings vectors
|
||||||
|
emb_result = self.embeddings(text=text)
|
||||||
|
vectors = emb_result.get("vectors", [])
|
||||||
|
|
||||||
# Query graph embeddings for semantic search
|
# Query graph embeddings for semantic search
|
||||||
input = {
|
input = {
|
||||||
"text": text,
|
"vectors": vectors,
|
||||||
"user": user,
|
"user": user,
|
||||||
"collection": collection,
|
"collection": collection,
|
||||||
"limit": limit
|
"limit": limit
|
||||||
|
|
@ -597,6 +601,51 @@ class FlowInstance:
|
||||||
input
|
input
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def document_embeddings_query(self, text, user, collection, limit=10):
|
||||||
|
"""
|
||||||
|
Query document chunks using semantic similarity.
|
||||||
|
|
||||||
|
Finds document chunks whose content is semantically similar to the
|
||||||
|
input text, using vector embeddings.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text: Query text for semantic search
|
||||||
|
user: User/keyspace identifier
|
||||||
|
collection: Collection identifier
|
||||||
|
limit: Maximum number of results (default: 10)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
dict: Query results with similar document chunks
|
||||||
|
|
||||||
|
Example:
|
||||||
|
```python
|
||||||
|
flow = api.flow().id("default")
|
||||||
|
results = flow.document_embeddings_query(
|
||||||
|
text="machine learning algorithms",
|
||||||
|
user="trustgraph",
|
||||||
|
collection="research-papers",
|
||||||
|
limit=5
|
||||||
|
)
|
||||||
|
```
|
||||||
|
"""
|
||||||
|
|
||||||
|
# First convert text to embeddings vectors
|
||||||
|
emb_result = self.embeddings(text=text)
|
||||||
|
vectors = emb_result.get("vectors", [])
|
||||||
|
|
||||||
|
# Query document embeddings for semantic search
|
||||||
|
input = {
|
||||||
|
"vectors": vectors,
|
||||||
|
"user": user,
|
||||||
|
"collection": collection,
|
||||||
|
"limit": limit
|
||||||
|
}
|
||||||
|
|
||||||
|
return self.request(
|
||||||
|
"service/document-embeddings",
|
||||||
|
input
|
||||||
|
)
|
||||||
|
|
||||||
def prompt(self, id, variables):
|
def prompt(self, id, variables):
|
||||||
"""
|
"""
|
||||||
Execute a prompt template with variable substitution.
|
Execute a prompt template with variable substitution.
|
||||||
|
|
|
||||||
|
|
@ -649,8 +649,12 @@ class SocketFlowInstance:
|
||||||
)
|
)
|
||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
|
# First convert text to embeddings vectors
|
||||||
|
emb_result = self.embeddings(text=text)
|
||||||
|
vectors = emb_result.get("vectors", [])
|
||||||
|
|
||||||
request = {
|
request = {
|
||||||
"text": text,
|
"vectors": vectors,
|
||||||
"user": user,
|
"user": user,
|
||||||
"collection": collection,
|
"collection": collection,
|
||||||
"limit": limit
|
"limit": limit
|
||||||
|
|
@ -659,6 +663,54 @@ class SocketFlowInstance:
|
||||||
|
|
||||||
return self.client._send_request_sync("graph-embeddings", self.flow_id, request, False)
|
return self.client._send_request_sync("graph-embeddings", self.flow_id, request, False)
|
||||||
|
|
||||||
|
def document_embeddings_query(
|
||||||
|
self,
|
||||||
|
text: str,
|
||||||
|
user: str,
|
||||||
|
collection: str,
|
||||||
|
limit: int = 10,
|
||||||
|
**kwargs: Any
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""
|
||||||
|
Query document chunks using semantic similarity.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text: Query text for semantic search
|
||||||
|
user: User/keyspace identifier
|
||||||
|
collection: Collection identifier
|
||||||
|
limit: Maximum number of results (default: 10)
|
||||||
|
**kwargs: Additional parameters passed to the service
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
dict: Query results with similar document chunks
|
||||||
|
|
||||||
|
Example:
|
||||||
|
```python
|
||||||
|
socket = api.socket()
|
||||||
|
flow = socket.flow("default")
|
||||||
|
|
||||||
|
results = flow.document_embeddings_query(
|
||||||
|
text="machine learning algorithms",
|
||||||
|
user="trustgraph",
|
||||||
|
collection="research-papers",
|
||||||
|
limit=5
|
||||||
|
)
|
||||||
|
```
|
||||||
|
"""
|
||||||
|
# First convert text to embeddings vectors
|
||||||
|
emb_result = self.embeddings(text=text)
|
||||||
|
vectors = emb_result.get("vectors", [])
|
||||||
|
|
||||||
|
request = {
|
||||||
|
"vectors": vectors,
|
||||||
|
"user": user,
|
||||||
|
"collection": collection,
|
||||||
|
"limit": limit
|
||||||
|
}
|
||||||
|
request.update(kwargs)
|
||||||
|
|
||||||
|
return self.client._send_request_sync("document-embeddings", self.flow_id, request, False)
|
||||||
|
|
||||||
def embeddings(self, text: str, **kwargs: Any) -> Dict[str, Any]:
|
def embeddings(self, text: str, **kwargs: Any) -> Dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
Generate vector embeddings for text.
|
Generate vector embeddings for text.
|
||||||
|
|
|
||||||
121
trustgraph-cli/trustgraph/cli/invoke_document_embeddings.py
Normal file
121
trustgraph-cli/trustgraph/cli/invoke_document_embeddings.py
Normal file
|
|
@ -0,0 +1,121 @@
|
||||||
|
"""
|
||||||
|
Queries document chunks by text similarity using vector embeddings.
|
||||||
|
Returns a list of matching document chunks, truncated to the specified length.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import os
|
||||||
|
from trustgraph.api import Api
|
||||||
|
|
||||||
|
default_url = os.getenv("TRUSTGRAPH_URL", 'http://localhost:8088/')
|
||||||
|
default_token = os.getenv("TRUSTGRAPH_TOKEN", None)
|
||||||
|
|
||||||
|
def truncate_chunk(chunk, max_length):
|
||||||
|
"""Truncate a chunk to max_length characters, adding ellipsis if needed."""
|
||||||
|
if len(chunk) <= max_length:
|
||||||
|
return chunk
|
||||||
|
return chunk[:max_length] + "..."
|
||||||
|
|
||||||
|
def query(url, flow_id, query_text, user, collection, limit, max_chunk_length, token=None):
|
||||||
|
|
||||||
|
# Create API client
|
||||||
|
api = Api(url=url, token=token)
|
||||||
|
socket = api.socket()
|
||||||
|
flow = socket.flow(flow_id)
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Call document embeddings query service
|
||||||
|
result = flow.document_embeddings_query(
|
||||||
|
text=query_text,
|
||||||
|
user=user,
|
||||||
|
collection=collection,
|
||||||
|
limit=limit
|
||||||
|
)
|
||||||
|
|
||||||
|
chunks = result.get("chunks", [])
|
||||||
|
for i, chunk in enumerate(chunks, 1):
|
||||||
|
truncated = truncate_chunk(chunk, max_chunk_length)
|
||||||
|
print(f"{i}. {truncated}")
|
||||||
|
|
||||||
|
finally:
|
||||||
|
# Clean up socket connection
|
||||||
|
socket.close()
|
||||||
|
|
||||||
|
def main():
|
||||||
|
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
prog='tg-invoke-document-embeddings',
|
||||||
|
description=__doc__,
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'-u', '--url',
|
||||||
|
default=default_url,
|
||||||
|
help=f'API URL (default: {default_url})',
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'-t', '--token',
|
||||||
|
default=default_token,
|
||||||
|
help='Authentication token (default: $TRUSTGRAPH_TOKEN)',
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'-f', '--flow-id',
|
||||||
|
default="default",
|
||||||
|
help=f'Flow ID (default: default)'
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'-U', '--user',
|
||||||
|
default="trustgraph",
|
||||||
|
help='User/keyspace (default: trustgraph)',
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'-c', '--collection',
|
||||||
|
default="default",
|
||||||
|
help='Collection (default: default)',
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'-l', '--limit',
|
||||||
|
type=int,
|
||||||
|
default=10,
|
||||||
|
help='Maximum number of results (default: 10)',
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'--max-chunk-length',
|
||||||
|
type=int,
|
||||||
|
default=200,
|
||||||
|
help='Truncate chunks to N characters (default: 200)',
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'query',
|
||||||
|
nargs=1,
|
||||||
|
help='Query text to search for similar document chunks',
|
||||||
|
)
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
try:
|
||||||
|
|
||||||
|
query(
|
||||||
|
url=args.url,
|
||||||
|
flow_id=args.flow_id,
|
||||||
|
query_text=args.query[0],
|
||||||
|
user=args.user,
|
||||||
|
collection=args.collection,
|
||||||
|
limit=args.limit,
|
||||||
|
max_chunk_length=args.max_chunk_length,
|
||||||
|
token=args.token,
|
||||||
|
)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
|
||||||
|
print("Exception:", e, flush=True)
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
77
trustgraph-cli/trustgraph/cli/invoke_embeddings.py
Normal file
77
trustgraph-cli/trustgraph/cli/invoke_embeddings.py
Normal file
|
|
@ -0,0 +1,77 @@
|
||||||
|
"""
|
||||||
|
Invokes the embeddings service to convert text to a vector embedding.
|
||||||
|
Returns the embedding vector as a list of floats.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import os
|
||||||
|
from trustgraph.api import Api
|
||||||
|
|
||||||
|
default_url = os.getenv("TRUSTGRAPH_URL", 'http://localhost:8088/')
|
||||||
|
default_token = os.getenv("TRUSTGRAPH_TOKEN", None)
|
||||||
|
|
||||||
|
def query(url, flow_id, text, token=None):
|
||||||
|
|
||||||
|
# Create API client
|
||||||
|
api = Api(url=url, token=token)
|
||||||
|
socket = api.socket()
|
||||||
|
flow = socket.flow(flow_id)
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Call embeddings service
|
||||||
|
result = flow.embeddings(text=text)
|
||||||
|
vectors = result.get("vectors", [])
|
||||||
|
print(vectors)
|
||||||
|
|
||||||
|
finally:
|
||||||
|
# Clean up socket connection
|
||||||
|
socket.close()
|
||||||
|
|
||||||
|
def main():
|
||||||
|
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
prog='tg-invoke-embeddings',
|
||||||
|
description=__doc__,
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'-u', '--url',
|
||||||
|
default=default_url,
|
||||||
|
help=f'API URL (default: {default_url})',
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'-t', '--token',
|
||||||
|
default=default_token,
|
||||||
|
help='Authentication token (default: $TRUSTGRAPH_TOKEN)',
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'-f', '--flow-id',
|
||||||
|
default="default",
|
||||||
|
help=f'Flow ID (default: default)'
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'text',
|
||||||
|
nargs=1,
|
||||||
|
help='Text to convert to embedding vector',
|
||||||
|
)
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
try:
|
||||||
|
|
||||||
|
query(
|
||||||
|
url=args.url,
|
||||||
|
flow_id=args.flow_id,
|
||||||
|
text=args.text[0],
|
||||||
|
token=args.token,
|
||||||
|
)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
|
||||||
|
print("Exception:", e, flush=True)
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
106
trustgraph-cli/trustgraph/cli/invoke_graph_embeddings.py
Normal file
106
trustgraph-cli/trustgraph/cli/invoke_graph_embeddings.py
Normal file
|
|
@ -0,0 +1,106 @@
|
||||||
|
"""
|
||||||
|
Queries graph entities by text similarity using vector embeddings.
|
||||||
|
Returns a list of matching graph entities.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import os
|
||||||
|
from trustgraph.api import Api
|
||||||
|
|
||||||
|
default_url = os.getenv("TRUSTGRAPH_URL", 'http://localhost:8088/')
|
||||||
|
default_token = os.getenv("TRUSTGRAPH_TOKEN", None)
|
||||||
|
|
||||||
|
def query(url, flow_id, query_text, user, collection, limit, token=None):
|
||||||
|
|
||||||
|
# Create API client
|
||||||
|
api = Api(url=url, token=token)
|
||||||
|
socket = api.socket()
|
||||||
|
flow = socket.flow(flow_id)
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Call graph embeddings query service
|
||||||
|
result = flow.graph_embeddings_query(
|
||||||
|
text=query_text,
|
||||||
|
user=user,
|
||||||
|
collection=collection,
|
||||||
|
limit=limit
|
||||||
|
)
|
||||||
|
|
||||||
|
entities = result.get("entities", [])
|
||||||
|
for entity in entities:
|
||||||
|
print(entity)
|
||||||
|
|
||||||
|
finally:
|
||||||
|
# Clean up socket connection
|
||||||
|
socket.close()
|
||||||
|
|
||||||
|
def main():
|
||||||
|
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
prog='tg-invoke-graph-embeddings',
|
||||||
|
description=__doc__,
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'-u', '--url',
|
||||||
|
default=default_url,
|
||||||
|
help=f'API URL (default: {default_url})',
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'-t', '--token',
|
||||||
|
default=default_token,
|
||||||
|
help='Authentication token (default: $TRUSTGRAPH_TOKEN)',
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'-f', '--flow-id',
|
||||||
|
default="default",
|
||||||
|
help=f'Flow ID (default: default)'
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'-U', '--user',
|
||||||
|
default="trustgraph",
|
||||||
|
help='User/keyspace (default: trustgraph)',
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'-c', '--collection',
|
||||||
|
default="default",
|
||||||
|
help='Collection (default: default)',
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'-l', '--limit',
|
||||||
|
type=int,
|
||||||
|
default=10,
|
||||||
|
help='Maximum number of results (default: 10)',
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'query',
|
||||||
|
nargs=1,
|
||||||
|
help='Query text to search for similar graph entities',
|
||||||
|
)
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
try:
|
||||||
|
|
||||||
|
query(
|
||||||
|
url=args.url,
|
||||||
|
flow_id=args.flow_id,
|
||||||
|
query_text=args.query[0],
|
||||||
|
user=args.user,
|
||||||
|
collection=args.collection,
|
||||||
|
limit=args.limit,
|
||||||
|
token=args.token,
|
||||||
|
)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
|
||||||
|
print("Exception:", e, flush=True)
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
Loading…
Add table
Add a link
Reference in a new issue