diff --git a/go/engine_native.go b/go/engine_native.go index 375c852..7e8cba9 100644 --- a/go/engine_native.go +++ b/go/engine_native.go @@ -339,6 +339,28 @@ func bufferResult(op string, call func(out, errBuf *C.CstxBuffer) C.CstxStatusCo return takeBuffer(&out), nil } +// protoResult is the protobuf member of this family. Structured data crosses +// the ABI as a generated message, so this is the shape most native calls need; +// the other helpers cover the scalar returns. +// +// T is the message and PT its pointer, which is what protobuf methods are +// declared on and what every caller wants back. Returning the message by value +// instead would copy its internal MessageState, mutex included. +func protoResult[T any, PT interface { + *T + proto.Message +}](op string, call func(out, errBuf *C.CstxBuffer) C.CstxStatusCode) (PT, error) { + data, err := bufferResult(op, call) + if err != nil { + return nil, err + } + message := PT(new(T)) + if err := proto.Unmarshal(data, message); err != nil { + return nil, fmt.Errorf("cstx: decode %T protobuf for %s: %w", message, op, err) + } + return message, nil +} + func textResult(op string, call func(out, errBuf *C.CstxBuffer) C.CstxStatusCode) (string, error) { data, err := bufferResult(op, call) if err != nil { @@ -375,17 +397,9 @@ func boolByte(value bool) uint8 { // --- runtime ------------------------------------------------------------- func (e *nativeEngine) lastChange(_ context.Context) (*cstxproto.GraphChangeSet, error) { - data, err := bufferResult("cstx.last_change", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.GraphChangeSet]("cstx.last_change", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { return C.cstx_last_change(e.handle, out, errBuf) }) - if err != nil { - return nil, err - } - var wire cstxproto.GraphChangeSet - if err := proto.Unmarshal(data, &wire); err != nil { - return nil, fmt.Errorf("cstx: decode change set protobuf: %w", err) - } - return &wire, nil } // --- extensions ---------------------------------------------------------- @@ -437,19 +451,11 @@ func (e *nativeEngine) extensionList(_ context.Context) (*cstxproto.ExtensionCat } func (e *nativeEngine) extensionInfo(_ context.Context, name string) (*cstxproto.ExtensionInfo, error) { - data, err := bufferResult("extensions.info", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.ExtensionInfo]("extensions.info", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { rc := C.cstx_extension_info(e.handle, stringSlice(name), out, errBuf) runtime.KeepAlive(name) return rc }) - if err != nil { - return nil, err - } - var value cstxproto.ExtensionInfo - if err := proto.Unmarshal(data, &value); err != nil { - return nil, err - } - return &value, nil } func (e *nativeEngine) extensionContains(_ context.Context, nodeType string) (bool, error) { @@ -515,71 +521,116 @@ func (e *nativeEngine) extensionAnchorConcepts(_ context.Context) (cstxproto.Anc // --- graph --------------------------------------------------------------- func (e *nativeEngine) graphParse(_ context.Context, payload *cstxproto.ParserPayload) (*cstxproto.Graph, uint64, error) { - graph, records, err := e.graphParseWire(context.Background(), payload) + if payload == nil { + return nil, 0, &Error{Code: CodeInvalidArgument, Operation: "graph.parse", Message: "payload must not be nil"} + } + encoded, err := proto.Marshal(payload) if err != nil { return nil, 0, err } - return &graph, records, nil + // Not protoResult: this is the only native call with a second out parameter. + var records C.uint64_t + var out, errBuf C.CstxBuffer + rc := C.cstx_graph_parse(e.handle, byteSlice(encoded), &records, &out, &errBuf) + runtime.KeepAlive(encoded) + if err := statusError(rc, "graph.parse", &errBuf); err != nil { + C.cstx_buffer_free(&out) + return nil, 0, err + } + graph := &cstxproto.Graph{} + if err := proto.Unmarshal(takeBuffer(&out), graph); err != nil { + return nil, 0, fmt.Errorf("cstx: decode graph protobuf for graph.parse: %w", err) + } + return graph, uint64(records), nil } func (e *nativeEngine) graphLink(_ context.Context, nodeIDs []string, dataSource string) (*cstxproto.GraphLinkResult, error) { - result, err := e.graphLinkWire(context.Background(), nodeIDs, dataSource) + selection, err := proto.Marshal(&cstxproto.GraphSelection{NodeIds: nodeIDs}) if err != nil { return nil, err } - return &result, nil + return protoResult[cstxproto.GraphLinkResult]("graph.link", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + rc := C.cstx_graph_link(e.handle, byteSlice(selection), stringSlice(dataSource), out, errBuf) + runtime.KeepAlive(selection) + runtime.KeepAlive(dataSource) + return rc + }) +} + +// graphWrite serializes one Graph aggregate and reports how many elements the +// native side actually changed. The three batch writes differ only in which C +// symbol they reach. +func (e *nativeEngine) graphWrite( + operation string, + graph *cstxproto.Graph, + call func(payload C.CstxSlice, out *C.uint64_t, errBuf *C.CstxBuffer) C.CstxStatusCode, +) (uint64, error) { + payload, err := proto.Marshal(graph) + if err != nil { + return 0, err + } + return countResult(operation, func(out *C.uint64_t, errBuf *C.CstxBuffer) C.CstxStatusCode { + rc := call(byteSlice(payload), out, errBuf) + runtime.KeepAlive(payload) + return rc + }) } func (e *nativeEngine) graphAddNodes(_ context.Context, nodes []*cstxproto.Node) (uint64, error) { - return e.graphAddNodesWire(context.Background(), &cstxproto.Graph{Nodes: nodes}) + return e.graphWrite("graph.add_nodes", &cstxproto.Graph{Nodes: nodes}, func(payload C.CstxSlice, out *C.uint64_t, errBuf *C.CstxBuffer) C.CstxStatusCode { + return C.cstx_graph_add_nodes(e.handle, payload, out, errBuf) + }) } func (e *nativeEngine) graphReplaceNodes(_ context.Context, nodes []*cstxproto.Node) (uint64, error) { - return e.graphReplaceNodesWire(context.Background(), &cstxproto.Graph{Nodes: nodes}) + return e.graphWrite("graph.replace_nodes", &cstxproto.Graph{Nodes: nodes}, func(payload C.CstxSlice, out *C.uint64_t, errBuf *C.CstxBuffer) C.CstxStatusCode { + return C.cstx_graph_replace_nodes(e.handle, payload, out, errBuf) + }) } func (e *nativeEngine) graphAddRelationships(_ context.Context, relationships []*cstxproto.Relationship) (uint64, error) { - return e.graphAddRelationshipsWire(context.Background(), &cstxproto.Graph{Relationships: relationships}) + return e.graphWrite("graph.add_relationships", &cstxproto.Graph{Relationships: relationships}, func(payload C.CstxSlice, out *C.uint64_t, errBuf *C.CstxBuffer) C.CstxStatusCode { + return C.cstx_graph_add_relationships(e.handle, payload, out, errBuf) + }) } func (e *nativeEngine) graphAddRelationship(_ context.Context, relationship *cstxproto.Relationship) (*cstxproto.Relationship, error) { - value, err := e.graphAddRelationshipWire(context.Background(), relationship) + if relationship == nil { + return nil, &Error{Code: CodeInvalidArgument, Operation: "graph.add_relationship", Message: "relationship must not be nil"} + } + payload, err := proto.Marshal(relationship) if err != nil { return nil, err } - return &value, nil + return protoResult[cstxproto.Relationship]("graph.add_relationship", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + rc := C.cstx_graph_add_relationship(e.handle, byteSlice(payload), out, errBuf) + runtime.KeepAlive(payload) + return rc + }) } func (e *nativeEngine) graphNode(_ context.Context, nodeID string) (*cstxproto.Node, error) { - node, err := e.graphNodeWire(context.Background(), nodeID) - if err != nil { - return nil, err - } - return &node, nil + return protoResult[cstxproto.Node]("graph.node", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + rc := C.cstx_graph_node(e.handle, stringSlice(nodeID), out, errBuf) + runtime.KeepAlive(nodeID) + return rc + }) } func (e *nativeEngine) graphRelationship(_ context.Context, relationshipID string) (*cstxproto.Relationship, error) { - relationship, err := e.graphRelationshipWire(context.Background(), relationshipID) - if err != nil { - return nil, err - } - return &relationship, nil + return protoResult[cstxproto.Relationship]("graph.relationship", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + rc := C.cstx_graph_relationship(e.handle, stringSlice(relationshipID), out, errBuf) + runtime.KeepAlive(relationshipID) + return rc + }) } func (e *nativeEngine) graphFindNode(_ context.Context, identifier string) (*cstxproto.Node, error) { - data, err := bufferResult("graph.find_node", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.Node]("graph.find_node", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { rc := C.cstx_graph_find_node(e.handle, stringSlice(identifier), out, errBuf) runtime.KeepAlive(identifier) return rc }) - if err != nil { - return nil, err - } - var node cstxproto.Node - if err := proto.Unmarshal(data, &node); err != nil { - return nil, fmt.Errorf("cstx: decode node protobuf: %w", err) - } - return &node, nil } func (e *nativeEngine) graphNodeTypes(_ context.Context) ([]string, error) { @@ -630,19 +681,11 @@ func (e *nativeEngine) graphPatchNodeAnnotations(_ context.Context, update *cstx } func (e *nativeEngine) graphFindAnchors(_ context.Context, conceptName string) (*cstxproto.GraphAnchorCatalog, error) { - data, err := bufferResult("graph.find_anchors", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.GraphAnchorCatalog]("graph.find_anchors", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { rc := C.cstx_graph_find_anchors(e.handle, stringSlice(conceptName), out, errBuf) runtime.KeepAlive(conceptName) return rc }) - if err != nil { - return nil, err - } - var catalog cstxproto.GraphAnchorCatalog - if err := proto.Unmarshal(data, &catalog); err != nil { - return nil, fmt.Errorf("cstx: decode graph anchor catalog protobuf: %w", err) - } - return &catalog, nil } func (e *nativeEngine) graphContains(_ context.Context, nodeID string) (bool, error) { @@ -666,17 +709,9 @@ func (e *nativeEngine) graphRelationshipCount(_ context.Context) (uint64, error) } func (e *nativeEngine) graphStats(_ context.Context) (*cstxproto.GraphStats, error) { - data, err := bufferResult("graph.stats", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.GraphStats]("graph.stats", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { return C.cstx_graph_stats(e.handle, 0, 0, out, errBuf) }) - if err != nil { - return nil, err - } - var wire cstxproto.GraphStats - if err := proto.Unmarshal(data, &wire); err != nil { - return nil, err - } - return &wire, nil } func (e *nativeEngine) graphNodes(_ context.Context, query *cstxproto.NodeQuery) (graphCursor, error) { @@ -797,19 +832,11 @@ func (e *nativeEngine) repoCheckout(_ context.Context, revision string, force bo if force { nativeForce = 1 } - data, err := bufferResult("repo.checkout", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.Commit]("repo.checkout", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { rc := C.cstx_repo_checkout(e.handle, stringSlice(revision), nativeForce, out, errBuf) runtime.KeepAlive(revision) return rc }) - if err != nil { - return nil, err - } - var wire cstxproto.Commit - if err := proto.Unmarshal(data, &wire); err != nil { - return nil, err - } - return &wire, nil } func (e *nativeEngine) repoCommit( @@ -826,7 +853,7 @@ func (e *nativeEngine) repoCommit( if err != nil { return nil, err } - data, err := bufferResult("repo.commit", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.Commit]("repo.commit", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { var expected C.CstxSlice if expectedHead != nil { expected = stringSlice(*expectedHead) @@ -838,14 +865,6 @@ func (e *nativeEngine) repoCommit( runtime.KeepAlive(metadataBytes) return rc }) - if err != nil { - return nil, err - } - var wire cstxproto.Commit - if err := proto.Unmarshal(data, &wire); err != nil { - return nil, err - } - return &wire, nil } func (e *nativeEngine) repoPrepare( @@ -869,7 +888,7 @@ func (e *nativeEngine) repoPrepare( nativeTimestamp = C.int64_t(*timestamp) hasTimestamp = 1 } - data, err := bufferResult("repo.prepare", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.PublicationPlan]("repo.prepare", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { var expected C.CstxSlice if expectedHead != nil { expected = stringSlice(*expectedHead) @@ -881,14 +900,6 @@ func (e *nativeEngine) repoPrepare( runtime.KeepAlive(metadataBytes) return rc }) - if err != nil { - return nil, err - } - var wire cstxproto.PublicationPlan - if err := proto.Unmarshal(data, &wire); err != nil { - return nil, err - } - return &wire, nil } func (e *nativeEngine) repoAccept(_ context.Context, commit string) error { @@ -936,19 +947,11 @@ func (e *nativeEngine) repoMissing(_ context.Context, plan *cstxproto.Repository if err != nil { return nil, err } - data, err := bufferResult("repo.missing", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.ObjectSelection]("repo.missing", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { rc := C.cstx_repo_missing(e.handle, byteSlice(payload), out, errBuf) runtime.KeepAlive(payload) return rc }) - if err != nil { - return nil, err - } - var value cstxproto.ObjectSelection - if err := proto.Unmarshal(data, &value); err != nil { - return nil, err - } - return &value, nil } func (e *nativeEngine) repoReleaseTransientObjects(_ context.Context) error { @@ -968,21 +971,13 @@ func (e *nativeEngine) repoDiff(_ context.Context, baseRef, headRef string, limi if detailValue == cstxproto.DiffDetail_DIFF_DETAIL_COUNTS { detail = "counts" } - data, err := bufferResult("repo.diff", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.GraphDiff]("repo.diff", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { rc := C.cstx_repo_diff(e.handle, stringSlice(baseRef), stringSlice(headRef), nativeLimit, hasLimit, stringSlice(detail), out, errBuf) runtime.KeepAlive(baseRef) runtime.KeepAlive(headRef) runtime.KeepAlive(detail) return rc }) - if err != nil { - return nil, err - } - var wire cstxproto.GraphDiff - if err := proto.Unmarshal(data, &wire); err != nil { - return nil, err - } - return &wire, nil } func (e *nativeEngine) repoHead(_ context.Context, refName string) (*string, error) { @@ -1001,19 +996,11 @@ func (e *nativeEngine) repoHead(_ context.Context, refName string) (*string, err } func (e *nativeEngine) repoLog(_ context.Context, revision string, limit int) (*cstxproto.CommitLog, error) { - data, err := bufferResult("repo.log", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.CommitLog]("repo.log", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { rc := C.cstx_repo_log(e.handle, stringSlice(revision), C.size_t(limit), out, errBuf) runtime.KeepAlive(revision) return rc }) - if err != nil { - return nil, err - } - var wire cstxproto.CommitLog - if err := proto.Unmarshal(data, &wire); err != nil { - return nil, err - } - return &wire, nil } func (e *nativeEngine) repoHistory( @@ -1028,20 +1015,12 @@ func (e *nativeEngine) repoHistory( nativeLimit = C.size_t(*limit) hasLimit = 1 } - data, err := bufferResult("repo.history", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.EntityHistory]("repo.history", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { rc := C.cstx_repo_history(e.handle, stringSlice(entityID), stringSlice(revision), nativeLimit, hasLimit, out, errBuf) runtime.KeepAlive(entityID) runtime.KeepAlive(revision) return rc }) - if err != nil { - return nil, err - } - var wire cstxproto.EntityHistory - if err := proto.Unmarshal(data, &wire); err != nil { - return nil, err - } - return &wire, nil } func (e *nativeEngine) repoEntities( @@ -1055,20 +1034,12 @@ func (e *nativeEngine) repoEntities( if err != nil { return nil, err } - data, err := bufferResult("repo.entities", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.Graph]("repo.entities", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { rc := C.cstx_repo_entities(e.handle, stringSlice(revision), byteSlice(selection), out, errBuf) runtime.KeepAlive(revision) runtime.KeepAlive(selection) return rc }) - if err != nil { - return nil, err - } - var wire cstxproto.Graph - if err := proto.Unmarshal(data, &wire); err != nil { - return nil, err - } - return &wire, nil } func (e *nativeEngine) repoBranch(_ context.Context, name, startPoint string) (string, error) { @@ -1087,7 +1058,7 @@ func (e *nativeEngine) repoMerge( expectedHead *string, message *string, ) (*cstxproto.Commit, error) { - data, err := bufferResult("repo.merge", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.Commit]("repo.merge", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { var expected, commitMessage C.CstxSlice if expectedHead != nil { expected = stringSlice(*expectedHead) @@ -1102,14 +1073,6 @@ func (e *nativeEngine) repoMerge( runtime.KeepAlive(message) return rc }) - if err != nil { - return nil, err - } - var wire cstxproto.Commit - if err := proto.Unmarshal(data, &wire); err != nil { - return nil, err - } - return &wire, nil } func (e *nativeEngine) repoStat( @@ -1118,19 +1081,11 @@ func (e *nativeEngine) repoStat( excludeMask uint64, includeMask uint64, ) (*cstxproto.GraphStats, error) { - data, err := bufferResult("repo.stat", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.GraphStats]("repo.stat", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { rc := C.cstx_repo_stat(e.handle, stringSlice(revision), C.uint64_t(excludeMask), C.uint64_t(includeMask), out, errBuf) runtime.KeepAlive(revision) return rc }) - if err != nil { - return nil, err - } - var wire cstxproto.GraphStats - if err := proto.Unmarshal(data, &wire); err != nil { - return nil, err - } - return &wire, nil } func (e *nativeEngine) repoDelta( @@ -1149,19 +1104,11 @@ func (e *nativeEngine) repoDelta( end = C.int64_t(*endTimestamp) hasEnd = 1 } - data, err := bufferResult("repo.delta", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.GraphChangeSummary]("repo.delta", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { rc := C.cstx_repo_delta(e.handle, stringSlice(revision), start, hasStart, end, hasEnd, out, errBuf) runtime.KeepAlive(revision) return rc }) - if err != nil { - return nil, err - } - var wire cstxproto.GraphChangeSummary - if err := proto.Unmarshal(data, &wire); err != nil { - return nil, err - } - return &wire, nil } // --- unified graph cursor ------------------------------------------------ @@ -1177,17 +1124,9 @@ func (c *nativeGraphCursor) page(_ context.Context, limit, page int) (*cstxproto if c.cursor == nil { return nil, &Error{Code: CodeInvalidArgument, Operation: "cursor.page", Message: "cursor is closed"} } - data, err := bufferResult("cursor.page", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.GraphResultPage]("cursor.page", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { return C.cstx_graph_cursor_page(c.cursor, C.size_t(limit), C.size_t(page), out, errBuf) }) - if err != nil { - return nil, err - } - var wire cstxproto.GraphResultPage - if err := proto.Unmarshal(data, &wire); err != nil { - return nil, err - } - return &wire, nil } func (c *nativeGraphCursor) close() { diff --git a/go/lib/darwin_amd64/libcstx_ffi.a b/go/lib/darwin_amd64/libcstx_ffi.a index 392c565..b86cf37 100644 Binary files a/go/lib/darwin_amd64/libcstx_ffi.a and b/go/lib/darwin_amd64/libcstx_ffi.a differ diff --git a/go/lib/darwin_arm64/libcstx_ffi.a b/go/lib/darwin_arm64/libcstx_ffi.a index c3c5d26..59b19ce 100644 Binary files a/go/lib/darwin_arm64/libcstx_ffi.a and b/go/lib/darwin_arm64/libcstx_ffi.a differ diff --git a/go/lib/linux_amd64/libcstx_ffi.a b/go/lib/linux_amd64/libcstx_ffi.a index 0f2b881..72e700b 100644 Binary files a/go/lib/linux_amd64/libcstx_ffi.a and b/go/lib/linux_amd64/libcstx_ffi.a differ diff --git a/go/lib/linux_arm64/libcstx_ffi.a b/go/lib/linux_arm64/libcstx_ffi.a index f28a3b4..6c0606a 100644 Binary files a/go/lib/linux_arm64/libcstx_ffi.a and b/go/lib/linux_arm64/libcstx_ffi.a differ diff --git a/go/lib/windows_amd64/libcstx_ffi.a b/go/lib/windows_amd64/libcstx_ffi.a index de07f25..0734603 100644 Binary files a/go/lib/windows_amd64/libcstx_ffi.a and b/go/lib/windows_amd64/libcstx_ffi.a differ diff --git a/go/proto_native.go b/go/proto_native.go deleted file mode 100644 index 3027eed..0000000 --- a/go/proto_native.go +++ /dev/null @@ -1,153 +0,0 @@ -package cstx - -// The native adapter transports only generated protobuf messages across the -// C ABI. It intentionally contains no SDK-owned graph or repository models. - -/* -#include "cstx_ffi.h" -*/ -import "C" - -import ( - "context" - "fmt" - "runtime" - - "github.com/chainreactors/libcstx/go/proto/cstxproto" - "google.golang.org/protobuf/proto" -) - -func (e *nativeEngine) graphParseWire(_ context.Context, payload *cstxproto.ParserPayload) (cstxproto.Graph, uint64, error) { - if payload == nil { - return cstxproto.Graph{}, 0, &Error{Code: CodeInvalidArgument, Operation: "graph.parse", Message: "payload must not be nil"} - } - encoded, err := proto.Marshal(payload) - if err != nil { - return cstxproto.Graph{}, 0, err - } - var records C.uint64_t - var out, errBuf C.CstxBuffer - if err := statusError(C.cstx_graph_parse(e.handle, byteSlice(encoded), &records, &out, &errBuf), "graph.parse", &errBuf); err != nil { - C.cstx_buffer_free(&out) - return cstxproto.Graph{}, 0, err - } - runtime.KeepAlive(encoded) - data := takeBuffer(&out) - var graph cstxproto.Graph - if err := proto.Unmarshal(data, &graph); err != nil { - return cstxproto.Graph{}, 0, fmt.Errorf("cstx: decode graph protobuf: %w", err) - } - return graph, uint64(records), nil -} - -func (e *nativeEngine) graphLinkWire(_ context.Context, nodeIDs []string, dataSource string) (cstxproto.GraphLinkResult, error) { - selection, err := proto.Marshal(&cstxproto.GraphSelection{NodeIds: nodeIDs}) - if err != nil { - return cstxproto.GraphLinkResult{}, err - } - data, err := bufferResult("graph.link", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { - rc := C.cstx_graph_link(e.handle, byteSlice(selection), stringSlice(dataSource), out, errBuf) - runtime.KeepAlive(selection) - runtime.KeepAlive(dataSource) - return rc - }) - if err != nil { - return cstxproto.GraphLinkResult{}, err - } - var result cstxproto.GraphLinkResult - if err := proto.Unmarshal(data, &result); err != nil { - return cstxproto.GraphLinkResult{}, fmt.Errorf("cstx: decode graph link result protobuf: %w", err) - } - return result, nil -} - -func (e *nativeEngine) graphAddNodesWire(_ context.Context, graph *cstxproto.Graph) (uint64, error) { - payload, err := proto.Marshal(graph) - if err != nil { - return 0, err - } - return countResult("graph.add_nodes", func(out *C.uint64_t, errBuf *C.CstxBuffer) C.CstxStatusCode { - rc := C.cstx_graph_add_nodes(e.handle, byteSlice(payload), out, errBuf) - runtime.KeepAlive(payload) - return rc - }) -} - -func (e *nativeEngine) graphReplaceNodesWire(_ context.Context, graph *cstxproto.Graph) (uint64, error) { - payload, err := proto.Marshal(graph) - if err != nil { - return 0, err - } - return countResult("graph.replace_nodes", func(out *C.uint64_t, errBuf *C.CstxBuffer) C.CstxStatusCode { - rc := C.cstx_graph_replace_nodes(e.handle, byteSlice(payload), out, errBuf) - runtime.KeepAlive(payload) - return rc - }) -} - -func (e *nativeEngine) graphAddRelationshipsWire(_ context.Context, graph *cstxproto.Graph) (uint64, error) { - payload, err := proto.Marshal(graph) - if err != nil { - return 0, err - } - return countResult("graph.add_relationships", func(out *C.uint64_t, errBuf *C.CstxBuffer) C.CstxStatusCode { - rc := C.cstx_graph_add_relationships(e.handle, byteSlice(payload), out, errBuf) - runtime.KeepAlive(payload) - return rc - }) -} - -func (e *nativeEngine) graphAddRelationshipWire(_ context.Context, relationship *cstxproto.Relationship) (cstxproto.Relationship, error) { - if relationship == nil { - return cstxproto.Relationship{}, &Error{Code: CodeInvalidArgument, Operation: "graph.add_relationship", Message: "relationship must not be nil"} - } - payload, err := proto.Marshal(relationship) - if err != nil { - return cstxproto.Relationship{}, err - } - data, err := bufferResult("graph.add_relationship", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { - rc := C.cstx_graph_add_relationship(e.handle, byteSlice(payload), out, errBuf) - runtime.KeepAlive(payload) - return rc - }) - if err != nil { - return cstxproto.Relationship{}, err - } - var stored cstxproto.Relationship - if err := proto.Unmarshal(data, &stored); err != nil { - return cstxproto.Relationship{}, fmt.Errorf("cstx: decode relationship protobuf: %w", err) - } - return stored, nil -} - -func (e *nativeEngine) graphNodeWire(_ context.Context, nodeID string) (cstxproto.Node, error) { - data, err := bufferResult("graph.node", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { - rc := C.cstx_graph_node(e.handle, stringSlice(nodeID), out, errBuf) - runtime.KeepAlive(nodeID) - return rc - }) - if err != nil { - return cstxproto.Node{}, err - } - var node cstxproto.Node - if err := proto.Unmarshal(data, &node); err != nil { - return cstxproto.Node{}, fmt.Errorf("cstx: decode node protobuf: %w", err) - } - return node, nil -} - -func (e *nativeEngine) graphRelationshipWire(_ context.Context, relationshipID string) (cstxproto.Relationship, error) { - data, err := bufferResult("graph.relationship", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { - rc := C.cstx_graph_relationship(e.handle, stringSlice(relationshipID), out, errBuf) - runtime.KeepAlive(relationshipID) - return rc - }) - if err != nil { - return cstxproto.Relationship{}, err - } - var relationship cstxproto.Relationship - if err := proto.Unmarshal(data, &relationship); err != nil { - return cstxproto.Relationship{}, fmt.Errorf("cstx: decode relationship protobuf: %w", err) - } - return relationship, nil -} diff --git a/go/rag_native.go b/go/rag_native.go index fca4e22..84935a5 100644 --- a/go/rag_native.go +++ b/go/rag_native.go @@ -71,34 +71,18 @@ func (s *nativeRagIndexSession) metadata(_ context.Context) (*cstxproto.RagIndex if s.session == nil { return nil, s.closedError("graph.rag.index.metadata") } - data, err := bufferResult("graph.rag.index.metadata", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.RagIndexResult]("graph.rag.index.metadata", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { return C.cstx_rag_index_session_metadata(s.session, out, errBuf) }) - if err != nil { - return nil, err - } - var result cstxproto.RagIndexResult - if err := proto.Unmarshal(data, &result); err != nil { - return nil, fmt.Errorf("cstx: decode rag index result protobuf: %w", err) - } - return &result, nil } func (s *nativeRagIndexSession) pending(_ context.Context, offset, limit int) (*cstxproto.RagRecordPage, error) { if s.session == nil { return nil, s.closedError("graph.rag.index.pending") } - data, err := bufferResult("graph.rag.index.pending", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.RagRecordPage]("graph.rag.index.pending", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { return C.cstx_rag_index_session_pending(s.session, C.size_t(offset), C.size_t(limit), out, errBuf) }) - if err != nil { - return nil, err - } - var page cstxproto.RagRecordPage - if err := proto.Unmarshal(data, &page); err != nil { - return nil, fmt.Errorf("cstx: decode rag record page protobuf: %w", err) - } - return &page, nil } func (s *nativeRagIndexSession) deletes(_ context.Context) ([]string, error) { @@ -198,17 +182,9 @@ func (r *nativeRagRetrieval) requests(_ context.Context) (*cstxproto.RecallPlan, if r.retrieval == nil { return nil, &Error{Code: CodeNotInitialized, Operation: "graph.rag.retrieve.requests", Message: "retrieval is closed"} } - data, err := bufferResult("graph.rag.retrieve.requests", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.RecallPlan]("graph.rag.retrieve.requests", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { return C.cstx_rag_retrieval_requests(r.retrieval, out, errBuf) }) - if err != nil { - return nil, err - } - var plan cstxproto.RecallPlan - if err := proto.Unmarshal(data, &plan); err != nil { - return nil, fmt.Errorf("cstx: decode recall plan protobuf: %w", err) - } - return &plan, nil } func (r *nativeRagRetrieval) complete(_ context.Context, results *cstxproto.RecallResults) (*cstxproto.RagResult, error) { @@ -219,19 +195,11 @@ func (r *nativeRagRetrieval) complete(_ context.Context, results *cstxproto.Reca if err != nil { return nil, err } - data, err := bufferResult("graph.rag.complete", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { + return protoResult[cstxproto.RagResult]("graph.rag.complete", func(out, errBuf *C.CstxBuffer) C.CstxStatusCode { rc := C.cstx_rag_retrieval_complete(r.retrieval, byteSlice(payload), out, errBuf) runtime.KeepAlive(payload) return rc }) - if err != nil { - return nil, err - } - var result cstxproto.RagResult - if err := proto.Unmarshal(data, &result); err != nil { - return nil, fmt.Errorf("cstx: decode rag result protobuf: %w", err) - } - return &result, nil } func (r *nativeRagRetrieval) close() { diff --git a/ts/wasm/cstx_wasm.d.ts b/ts/wasm/cstx_wasm.d.ts index fbc2776..e12663c 100644 --- a/ts/wasm/cstx_wasm.d.ts +++ b/ts/wasm/cstx_wasm.d.ts @@ -54,6 +54,10 @@ export class Graph { nodeCount(): bigint; nodeTypes(): any; nodes(options?: any | null): GraphCursor; + /** + * Parse one extension artifact without mutating this graph. + */ + parse(plugin: string, artifact: string, data: Uint8Array): any; patchNodeAnnotations(node_ids: string[] | null | undefined, patch: any): bigint; query(expression: string, options?: any | null): GraphCursor; querySubgraph(expression: string, limit?: number | null, page?: number | null, exclude_mask?: bigint | null, include_mask?: bigint | null): CSTX; @@ -118,6 +122,7 @@ export interface InitOutput { readonly extensions_schemas: (a: number, b: number) => void; readonly extensions_parsesArtifact: (a: number, b: number, c: number, d: number) => void; readonly extensions_anchorConcepts: (a: number, b: number) => void; + readonly graph_parse: (a: number, b: number, c: number, d: number, e: number, f: number, g: number, h: number) => void; readonly graph_addNode: (a: number, b: number, c: number) => void; readonly graph_addNodes: (a: number, b: number, c: number) => void; readonly graph_addRelationships: (a: number, b: number, c: number) => void; diff --git a/ts/wasm/cstx_wasm.js b/ts/wasm/cstx_wasm.js index 4c6ad1b..0ff0b81 100644 --- a/ts/wasm/cstx_wasm.js +++ b/ts/wasm/cstx_wasm.js @@ -717,6 +717,34 @@ export class Graph { wasm.__wbindgen_add_to_stack_pointer(16); } } + /** + * Parse one extension artifact without mutating this graph. + * @param {string} plugin + * @param {string} artifact + * @param {Uint8Array} data + * @returns {any} + */ + parse(plugin, artifact, data) { + try { + const retptr = wasm.__wbindgen_add_to_stack_pointer(-16); + const ptr0 = passStringToWasm0(plugin, wasm.__wbindgen_export, wasm.__wbindgen_export2); + const len0 = WASM_VECTOR_LEN; + const ptr1 = passStringToWasm0(artifact, wasm.__wbindgen_export, wasm.__wbindgen_export2); + const len1 = WASM_VECTOR_LEN; + const ptr2 = passArray8ToWasm0(data, wasm.__wbindgen_export); + const len2 = WASM_VECTOR_LEN; + wasm.graph_parse(retptr, this.__wbg_ptr, ptr0, len0, ptr1, len1, ptr2, len2); + var r0 = getDataViewMemory0().getInt32(retptr + 4 * 0, true); + var r1 = getDataViewMemory0().getInt32(retptr + 4 * 1, true); + var r2 = getDataViewMemory0().getInt32(retptr + 4 * 2, true); + if (r2) { + throw takeObject(r1); + } + return takeObject(r0); + } finally { + wasm.__wbindgen_add_to_stack_pointer(16); + } + } /** * @param {string[] | null | undefined} node_ids * @param {any} patch @@ -1750,6 +1778,13 @@ function isLikeNone(x) { return x === undefined || x === null; } +function passArray8ToWasm0(arg, malloc) { + const ptr = malloc(arg.length * 1, 1) >>> 0; + getUint8ArrayMemory0().set(arg, ptr / 1); + WASM_VECTOR_LEN = arg.length; + return ptr; +} + function passArrayJsValueToWasm0(array, malloc) { const ptr = malloc(array.length * 4, 4) >>> 0; const mem = getDataViewMemory0(); diff --git a/ts/wasm/cstx_wasm_bg.wasm b/ts/wasm/cstx_wasm_bg.wasm index 2d35604..4cb5d3e 100644 Binary files a/ts/wasm/cstx_wasm_bg.wasm and b/ts/wasm/cstx_wasm_bg.wasm differ diff --git a/ts/wasm/cstx_wasm_bg.wasm.d.ts b/ts/wasm/cstx_wasm_bg.wasm.d.ts index 0d6a70b..a3bcf14 100644 --- a/ts/wasm/cstx_wasm_bg.wasm.d.ts +++ b/ts/wasm/cstx_wasm_bg.wasm.d.ts @@ -16,6 +16,7 @@ export const extensions_schema: (a: number, b: number, c: number, d: number) => export const extensions_schemas: (a: number, b: number) => void; export const extensions_parsesArtifact: (a: number, b: number, c: number, d: number) => void; export const extensions_anchorConcepts: (a: number, b: number) => void; +export const graph_parse: (a: number, b: number, c: number, d: number, e: number, f: number, g: number, h: number) => void; export const graph_addNode: (a: number, b: number, c: number) => void; export const graph_addNodes: (a: number, b: number, c: number) => void; export const graph_addRelationships: (a: number, b: number, c: number) => void;