OSVauco/agents/rag/setup_corpus.py
chrischristiansen-glitch a2790be421 fix: create corpus via REST API directly, bypassing SDK backend_config bug
SDK 1.153.1 always defaults to Spanner when backend_config is omitted,
and crashes when it's provided. Use REST POST directly with
vectorDbConfig.ragManagedDb set to force serverless.
2026-05-24 19:11:25 +02:00

201 lines
6.3 KiB
Python

#!/usr/bin/env python3
"""
setup_corpus.py — Create a Vertex AI RAG Engine corpus and import documents.
Project: propane-will-491900-m5
SDK 1.153.1 has a bug where backend_config crashes when provided,
and defaults to Spanner when omitted. We bypass it by calling the
REST API directly for corpus creation, then use the SDK for everything else.
"""
import os
import json
import time
import subprocess
import vertexai
from vertexai import rag
from vertexai.rag import RagCorpus
PROJECT_ID = "propane-will-491900-m5"
LOCATION = os.environ.get("RAG_LOCATION", "us-central1")
CORPUS_DISPLAY_NAME = os.environ.get("RAG_CORPUS_NAME", "osvauco-knowledge-base")
GCS_SOURCE = os.environ.get(
"RAG_GCS_SOURCE",
f"gs://{PROJECT_ID}-agent-staging/rag-docs/",
)
def get_token() -> str:
return subprocess.check_output(
["gcloud", "auth", "print-access-token"], text=True
).strip()
def ensure_serverless_engine_config() -> None:
"""
Set project-level RAG Engine Config to basic (serverless) tier.
Polls the operation until done.
"""
print("Ensuring RAG Engine Config is set to serverless (basic) tier...")
endpoint = (
f"https://{LOCATION}-aiplatform.googleapis.com/v1beta1"
f"/projects/{PROJECT_ID}/locations/{LOCATION}/ragEngineConfig"
)
payload = json.dumps({"ragManagedDbConfig": {"basic": {}}})
token = get_token()
result = subprocess.run(
["curl", "-s", "-X", "PATCH",
"-H", f"Authorization: Bearer {token}",
"-H", "Content-Type: application/json",
endpoint, "-d", payload],
capture_output=True, text=True,
)
resp = json.loads(result.stdout)
if "error" in resp:
err = resp["error"]
print(f" ⚠️ Engine config warning ({err.get('code')}): {err.get('message')} — continuing.")
return
op_name = resp.get("name", "")
if "/operations/" in op_name and not resp.get("done"):
print(f" Polling operation {op_name.split('/')[-1]}...")
op_url = f"https://{LOCATION}-aiplatform.googleapis.com/v1beta1/{op_name}"
for _ in range(20):
time.sleep(3)
token = get_token()
r = subprocess.run(
["curl", "-s", "-H", f"Authorization: Bearer {token}", op_url],
capture_output=True, text=True,
)
op = json.loads(r.stdout)
if op.get("done"):
break
else:
print(" ⚠️ Operation timed out — continuing anyway.")
return
print(" ✓ RAG Engine Config set to basic tier.")
def create_corpus_rest() -> str:
"""
Create corpus via REST API directly, bypassing SDK backend_config bug.
Returns the corpus resource name.
"""
url = (
f"https://{LOCATION}-aiplatform.googleapis.com/v1beta1"
f"/projects/{PROJECT_ID}/locations/{LOCATION}/ragCorpora"
)
payload = json.dumps({
"displayName": CORPUS_DISPLAY_NAME,
"ragEmbeddingModelConfig": {
"vertexPredictionEndpoint": {
"model": "publishers/google/models/text-embedding-004"
}
},
"vectorDbConfig": {
"ragManagedDb": {}
}
})
token = get_token()
result = subprocess.run(
["curl", "-s", "-X", "POST",
"-H", f"Authorization: Bearer {token}",
"-H", "Content-Type: application/json",
url, "-d", payload],
capture_output=True, text=True,
)
resp = json.loads(result.stdout)
if "error" in resp:
raise RuntimeError(f"Failed to create corpus: {resp['error']}")
# REST returns a long-running operation
op_name = resp.get("name", "")
if "/operations/" not in op_name:
raise RuntimeError(f"Unexpected response: {resp}")
print(f" Polling corpus creation operation...")
op_url = f"https://{LOCATION}-aiplatform.googleapis.com/v1beta1/{op_name}"
for _ in range(40):
time.sleep(5)
token = get_token()
r = subprocess.run(
["curl", "-s", "-H", f"Authorization: Bearer {token}", op_url],
capture_output=True, text=True,
)
op = json.loads(r.stdout)
if op.get("done"):
if "error" in op:
raise RuntimeError(f"Corpus creation failed: {op['error']}")
corpus_name = op["response"]["name"]
return corpus_name
raise RuntimeError("Corpus creation operation timed out.")
def get_or_create_corpus() -> RagCorpus:
"""Return existing corpus by display name, or create a new serverless one."""
for c in rag.list_corpora():
if c.display_name == CORPUS_DISPLAY_NAME:
print(f"Corpus '{CORPUS_DISPLAY_NAME}' already exists: {c.name}")
return c
print(f"Creating corpus '{CORPUS_DISPLAY_NAME}' in {LOCATION} via REST...")
corpus_name = create_corpus_rest()
print(f"✓ Corpus created: {corpus_name}")
# Re-fetch via SDK so we get a proper RagCorpus object
for c in rag.list_corpora():
if c.name == corpus_name:
return c
raise RuntimeError(f"Corpus created but not found in list: {corpus_name}")
def import_documents(corpus: RagCorpus) -> None:
print(f"Importing files from {GCS_SOURCE}...")
rag.import_files(
corpus.name,
paths=[GCS_SOURCE],
transformation_config=rag.TransformationConfig(
chunking_config=rag.ChunkingConfig(
chunk_size=512,
chunk_overlap=100,
)
),
)
print("✓ Import complete.")
def test_retrieval(corpus: RagCorpus) -> None:
print("Running test retrieval query...")
response = rag.retrieval_query(
rag_resources=[rag.RagResource(rag_corpus=corpus.name)],
text="test query",
rag_retrieval_config=rag.RagRetrievalConfig(top_k=3),
)
print(f"✓ Test retrieval returned {len(response.contexts.contexts)} chunk(s).")
def main() -> None:
vertexai.init(project=PROJECT_ID, location=LOCATION)
ensure_serverless_engine_config()
corpus = get_or_create_corpus()
import_documents(corpus)
print(f"\nRAG_CORPUS={corpus.name}")
print("Add this to Secret Manager:")
print(f" gcloud secrets create rag-corpus-name --data-file=- <<<'{corpus.name}'")
test_retrieval(corpus)
if __name__ == "__main__":
main()