diff --git a/src/client/acontext-py/src/acontext/resources/async_sessions.py b/src/client/acontext-py/src/acontext/resources/async_sessions.py index ca28155dd..3c4ea7712 100644 --- a/src/client/acontext-py/src/acontext/resources/async_sessions.py +++ b/src/client/acontext-py/src/acontext/resources/async_sessions.py @@ -16,6 +16,7 @@ Message, MessageObservingStatus, Session, + SessionSearchResult, TokenCounts, ) from ..uploads import FileUpload, normalize_file_upload @@ -502,3 +503,31 @@ async def patch_configs( json_data=payload, ) return data.get("configs", {}) # type: ignore + + async def search( + self, + *, + query: str, + user_id: str, + limit: int | None = None, + ) -> SessionSearchResult: + """Search for sessions by semantic similarity to a query string. + + Args: + query: The search query text. + user_id: The User ID to search within. + limit: Maximum number of results to return (1-100, default 10). + + Returns: + SessionSearchResult containing list of matching session UUIDs. + + Example: + >>> result = await client.sessions.search(query="conversations", user_id="user_123") + >>> for session_id in result.session_ids: + ... print(session_id) + """ + params = build_params(query=query, user_id=user_id, limit=limit) + data = await self._requester.request( + "GET", "/sessions/search", params=params or None + ) + return SessionSearchResult.model_validate(data) diff --git a/src/client/acontext-py/src/acontext/resources/sessions.py b/src/client/acontext-py/src/acontext/resources/sessions.py index 5c15081c3..cadbf40fe 100644 --- a/src/client/acontext-py/src/acontext/resources/sessions.py +++ b/src/client/acontext-py/src/acontext/resources/sessions.py @@ -16,6 +16,7 @@ Message, MessageObservingStatus, Session, + SessionSearchResult, TokenCounts, ) from ..uploads import FileUpload, normalize_file_upload @@ -496,3 +497,35 @@ def patch_configs( json_data=payload, ) return data.get("configs", {}) # type: ignore + + def search( + self, + *, + query: str, + user_id: str, + project_id: str, + limit: int | None = None, + ) -> SessionSearchResult: + """Search for sessions by semantic similarity to a query string. + + Args: + query: The search query text. + user_id: The User ID to search within. + project_id: The Project ID to search within. + limit: Maximum number of results to return (1-100, default 10). + + Returns: + SessionSearchResult containing list of matching session UUIDs. + + Example: + >>> result = client.sessions.search( + ... query="conversations about authentication", + ... user_id="user_123", + ... project_id="proj_456" + ... ) + >>> for session_id in result.session_ids: + ... print(session_id) + """ + params = build_params(query=query, user_id=user_id, project_id=project_id, limit=limit) + data = self._requester.request("GET", "/sessions/search", params=params or None) + return SessionSearchResult.model_validate(data) diff --git a/src/client/acontext-py/src/acontext/types/session.py b/src/client/acontext-py/src/acontext/types/session.py index 400da55b4..03b8a8280 100644 --- a/src/client/acontext-py/src/acontext/types/session.py +++ b/src/client/acontext-py/src/acontext/types/session.py @@ -303,3 +303,9 @@ class MessageObservingStatus(BaseModel): ) pending: int = Field(..., description="Number of messages with pending status") updated_at: str = Field(..., description="Timestamp when the status was retrieved") + + +class SessionSearchResult(BaseModel): + """Response model for session search.""" + + session_ids: list[str] = Field(..., description="List of matching session UUIDs") diff --git a/src/client/acontext-ts/src/resources/sessions.ts b/src/client/acontext-ts/src/resources/sessions.ts index dd3cc4835..a620702bc 100644 --- a/src/client/acontext-ts/src/resources/sessions.ts +++ b/src/client/acontext-ts/src/resources/sessions.ts @@ -438,4 +438,40 @@ export class SessionsAPI { }); return (data as { configs: Record }).configs ?? {}; } + + /** + * Search for sessions by semantic similarity to a query string. + * + * @param options - Options for searching sessions. + * @param options.query - The search query text. + * @param options.userId - The User ID to search within. + * @param options.projectId - The Project ID to search within. + * @param options.limit - Maximum number of results to return (1-100, default 10). + * @returns SessionSearchResult containing list of matching session UUIDs. + * + * @example + * const result = await client.sessions.search({ + * query: 'conversations about authentication', + * userId: 'user_123', + * projectId: 'proj_456' + * }); + */ + async search(options: { + query: string; + userId: string; + projectId: string; + limit?: number | null; + }): Promise<{ session_ids: string[] }> { + const params = buildParams({ + query: options.query, + user_id: options.userId, + project_id: options.projectId, + limit: options.limit ?? null, + }); + const data = await this.requester.request('GET', '/sessions/search', { + params: Object.keys(params).length > 0 ? params : undefined, + }); + // Assuming SessionSearchResult schema exists or returning generic object + return data as { session_ids: string[] }; + } } diff --git a/src/server/api/go/configs/config.yaml b/src/server/api/go/configs/config.yaml index 1016e6617..62f8a3ceb 100644 --- a/src/server/api/go/configs/config.yaml +++ b/src/server/api/go/configs/config.yaml @@ -52,3 +52,6 @@ telemetry: artifact: maxUploadSizeBytes: ${ARTIFACT_MAX_UPLOAD_SIZE_BYTES} # Default 16MB (16 * 1024 * 1024 bytes) + +embedding: + taskVectorDim: ${TASK_VECTOR_DIMENSION} diff --git a/src/server/api/go/go.mod b/src/server/api/go/go.mod index ff47a316c..e118ce0d9 100644 --- a/src/server/api/go/go.mod +++ b/src/server/api/go/go.mod @@ -15,6 +15,7 @@ require ( github.com/go-playground/validator/v10 v10.30.1 github.com/google/uuid v1.6.0 github.com/openai/openai-go/v3 v3.22.0 + github.com/pgvector/pgvector-go v0.2.2 github.com/rabbitmq/amqp091-go v1.10.0 github.com/redis/go-redis/extra/redisotel/v9 v9.18.0 github.com/redis/go-redis/v9 v9.18.0 diff --git a/src/server/api/go/go.sum b/src/server/api/go/go.sum index 425422f80..f44d16d99 100644 --- a/src/server/api/go/go.sum +++ b/src/server/api/go/go.sum @@ -4,6 +4,8 @@ cloud.google.com/go/auth v0.18.1 h1:IwTEx92GFUo2pJ6Qea0EU3zYvKnTAeRCODxfA/G5UWs= cloud.google.com/go/auth v0.18.1/go.mod h1:GfTYoS9G3CWpRA3Va9doKN9mjPGRS+v41jmZAhBzbrA= cloud.google.com/go/compute/metadata v0.9.0 h1:pDUj4QMoPejqq20dK0Pg2N4yG9zIkYGdBtwLoEkH9Zs= cloud.google.com/go/compute/metadata v0.9.0/go.mod h1:E0bWwX5wTnLPedCKqk3pJmVgCBSM6qQI1yTBdEb3C10= +entgo.io/ent v0.13.1 h1:uD8QwN1h6SNphdCCzmkMN3feSUzNnVvV/WIkHKMbzOE= +entgo.io/ent v0.13.1/go.mod h1:qCEmo+biw3ccBn9OyL4ZK5dfpwg++l1Gxwac5B1206A= filippo.io/edwards25519 v1.1.0 h1:FNf4tywRC1HmFuKW5xopWpigGjJKiJSV0Cqo0cJWDaA= filippo.io/edwards25519 v1.1.0/go.mod h1:BxyFTGdWcka3PhytdK4V28tE5sGfRvvvRV7EaN4VDT4= github.com/ClickHouse/ch-go v0.70.0 h1:/0lJpiSXxg/7IaJi7TOkKAOHrx0z0OiSMU475EJNAwM= @@ -139,6 +141,10 @@ github.com/go-openapi/testify/enable/yaml/v2 v2.0.2 h1:0+Y41Pz1NkbTHz8NngxTuAXxE github.com/go-openapi/testify/enable/yaml/v2 v2.0.2/go.mod h1:kme83333GCtJQHXQ8UKX3IBZu6z8T5Dvy5+CW3NLUUg= github.com/go-openapi/testify/v2 v2.0.2 h1:X999g3jeLcoY8qctY/c/Z8iBHTbwLz7R2WXd6Ub6wls= github.com/go-openapi/testify/v2 v2.0.2/go.mod h1:HCPmvFFnheKK2BuwSA0TbbdxJ3I16pjwMkYkP4Ywn54= +github.com/go-pg/pg/v10 v10.11.0 h1:CMKJqLgTrfpE/aOVeLdybezR2om071Vh38OLZjsyMI0= +github.com/go-pg/pg/v10 v10.11.0/go.mod h1:4BpHRoxE61y4Onpof3x1a2SQvi9c+q1dJnrNdMjsroA= +github.com/go-pg/zerochecker v0.2.0 h1:pp7f72c3DobMWOb2ErtZsnrPaSvHd2W4o9//8HtF4mU= +github.com/go-pg/zerochecker v0.2.0/go.mod h1:NJZ4wKL0NmTtz0GKCoJ8kym6Xn/EQzXRl2OnAe7MmDo= github.com/go-playground/assert/v2 v2.2.0 h1:JvknZsQTYeFEAhQwI4qEt9cyV5ONwRHC+lYKSsYSR8s= github.com/go-playground/assert/v2 v2.2.0/go.mod h1:VDjEfimB/XKnb+ZQfWdccd7VUvScMdVu0Titje2rxJ4= github.com/go-playground/locales v0.14.1 h1:EWaQ/wswjilfKLTECiXz7Rh+3BjFhfDFKv/oXslEjJA= @@ -195,6 +201,8 @@ github.com/jinzhu/inflection v1.0.0 h1:K317FqzuhWc8YvSVlFMCCUb36O/S9MCKRDI7QkRKD github.com/jinzhu/inflection v1.0.0/go.mod h1:h+uFLlag+Qp1Va5pdKtLDYj+kHp5pxUVkryuEj+Srlc= github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ= github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8= +github.com/jmoiron/sqlx v1.3.5 h1:vFFPA71p1o5gAeqtEAwLU4dnX2napprKtHr7PYIcN3g= +github.com/jmoiron/sqlx v1.3.5/go.mod h1:nRVWtLre0KfCLJvgxzCsLVMogSvQ1zNJtpYr2Ccp0mQ= github.com/json-iterator/go v1.1.12 h1:PV8peI4a0ysnczrg+LtxykD8LfKY9ML6u2jnxaEnrnM= github.com/json-iterator/go v1.1.12/go.mod h1:e30LSqwooZae/UwlEbR2852Gd8hjQvJoHmT4TnhNGBo= github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8= @@ -213,6 +221,8 @@ github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= github.com/leodido/go-urn v1.4.0 h1:WT9HwE9SGECu3lg4d/dIA+jxlljEa1/ffXKmRjqdmIQ= github.com/leodido/go-urn v1.4.0/go.mod h1:bvxc+MVxLKB4z00jd1z+Dvzr47oO32F/QSNjSBOlFxI= +github.com/lib/pq v1.10.9 h1:YXG7RB+JIjhP29X+OtkiDnYaXQwpS4JEWq7dtCCRUEw= +github.com/lib/pq v1.10.9/go.mod h1:AlVN5x4E4T544tWzH6hKfbfQvm3HdbOxrmggDNAPY9o= github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= github.com/mattn/go-sqlite3 v1.14.22 h1:2gZY6PC6kBnID23Tichd1K+Z0oS6nE/XwU+Vz/5o4kU= @@ -232,6 +242,8 @@ github.com/paulmach/orb v0.12.0/go.mod h1:5mULz1xQfs3bmQm63QEJA6lNGujuRafwA5S/En github.com/paulmach/protoscan v0.2.1/go.mod h1:SpcSwydNLrxUGSDvXvO0P7g7AuhJ7lcKfDlhJCDw2gY= github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4= github.com/pelletier/go-toml/v2 v2.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY= +github.com/pgvector/pgvector-go v0.2.2 h1:Q/oArmzgbEcio88q0tWQksv/u9Gnb1c3F1K2TnalxR0= +github.com/pgvector/pgvector-go v0.2.2/go.mod h1:u5sg3z9bnqVEdpe1pkTij8/rFhTaMCMNyQagPDLK8gQ= github.com/pierrec/lz4/v4 v4.1.25 h1:kocOqRffaIbU5djlIBr7Wh+cx82C0vtFb0fOurZHqD0= github.com/pierrec/lz4/v4 v4.1.25/go.mod h1:EoQMVJgeeEOMsCqCzqFm2O0cJvljX2nGZjcRIPL34O4= github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= @@ -303,10 +315,26 @@ github.com/tidwall/sjson v1.2.5 h1:kLy8mja+1c9jlljvWTlSazM7cKDRfJuR/bOJhcY5NcY= github.com/tidwall/sjson v1.2.5/go.mod h1:Fvgq9kS/6ociJEDnK0Fk1cpYF4FIW6ZF7LAe+6jwd28= github.com/tiktoken-go/tokenizer v0.7.0 h1:VMu6MPT0bXFDHr7UPh9uii7CNItVt3X9K90omxL54vw= github.com/tiktoken-go/tokenizer v0.7.0/go.mod h1:6UCYI/DtOallbmL7sSy30p6YQv60qNyU/4aVigPOx6w= +github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc h1:9lRDQMhESg+zvGYmW5DyG0UqvY96Bu5QYsTLvCHdrgo= +github.com/tmthrgd/go-hex v0.0.0-20190904060850-447a3041c3bc/go.mod h1:bciPuU6GHm1iF1pBvUfxfsH0Wmnc2VbpgvbI9ZWuIRs= github.com/twitchyliquid64/golang-asm v0.15.1 h1:SU5vSMR7hnwNxj24w34ZyCi/FmDZTkS4MhqMhdFk5YI= github.com/twitchyliquid64/golang-asm v0.15.1/go.mod h1:a1lVb/DtPvCB8fslRZhAngC2+aY1QWCk3Cedj/Gdt08= github.com/ugorji/go/codec v1.3.1 h1:waO7eEiFDwidsBN6agj1vJQ4AG7lh2yqXyOXqhgQuyY= github.com/ugorji/go/codec v1.3.1/go.mod h1:pRBVtBSKl77K30Bv8R2P+cLSGaTtex6fsA2Wjqmfxj4= +github.com/uptrace/bun v1.1.12 h1:sOjDVHxNTuM6dNGaba0wUuz7KvDE1BmNu9Gqs2gJSXQ= +github.com/uptrace/bun v1.1.12/go.mod h1:NPG6JGULBeQ9IU6yHp7YGELRa5Agmd7ATZdz4tGZ6z0= +github.com/uptrace/bun/dialect/pgdialect v1.1.12 h1:m/CM1UfOkoBTglGO5CUTKnIKKOApOYxkcP2qn0F9tJk= +github.com/uptrace/bun/dialect/pgdialect v1.1.12/go.mod h1:Ij6WIxQILxLlL2frUBxUBOZJtLElD2QQNDcu/PWDHTc= +github.com/uptrace/bun/driver/pgdriver v1.1.12 h1:3rRWB1GK0psTJrHwxzNfEij2MLibggiLdTqjTtfHc1w= +github.com/uptrace/bun/driver/pgdriver v1.1.12/go.mod h1:ssYUP+qwSEgeDDS1xm2XBip9el1y9Mi5mTAvLoiADLM= +github.com/vmihailenco/bufpool v0.1.11 h1:gOq2WmBrq0i2yW5QJ16ykccQ4wH9UyEsgLm6czKAd94= +github.com/vmihailenco/bufpool v0.1.11/go.mod h1:AFf/MOy3l2CFTKbxwt0mp2MwnqjNEs5H/UxrkA5jxTQ= +github.com/vmihailenco/msgpack/v5 v5.3.5 h1:5gO0H1iULLWGhs2H5tbAHIZTV8/cYafcFOr9znI5mJU= +github.com/vmihailenco/msgpack/v5 v5.3.5/go.mod h1:7xyJ9e+0+9SaZT0Wt1RGleJXzli6Q/V5KbhBonMG9jc= +github.com/vmihailenco/tagparser v0.1.2 h1:gnjoVuB/kljJ5wICEEOpx98oXMWPLj22G67Vbd1qPqc= +github.com/vmihailenco/tagparser v0.1.2/go.mod h1:OeAg3pn3UbLjkWt+rN9oFYB6u/cQgqMEUPoW2WPyhdI= +github.com/vmihailenco/tagparser/v2 v2.0.0 h1:y09buUbR+b5aycVFQs/g70pqKVZNBmxwAhO7/IwNM9g= +github.com/vmihailenco/tagparser/v2 v2.0.0/go.mod h1:Wri+At7QHww0WTrCBeu4J6bNtoV6mEfg5OIWRZA9qds= github.com/xdg-go/pbkdf2 v1.0.0/go.mod h1:jrpuAogTd400dnrH08LKmI/xc1MbPOebTwRqcT5RDeI= github.com/xdg-go/scram v1.1.1/go.mod h1:RaEWvsqvNKKvBPvcKeFjrG2cJqOkHTiyTpzz23ni57g= github.com/xdg-go/stringprep v1.0.3/go.mod h1:W3f5j4i+9rC0kuIEJL0ky1VpHXQU3ocBgklLGvcBnW8= @@ -460,3 +488,5 @@ gorm.io/gorm v1.31.1 h1:7CA8FTFz/gRfgqgpeKIBcervUn3xSyPUmr6B2WXJ7kg= gorm.io/gorm v1.31.1/go.mod h1:XyQVbO2k6YkOis7C2437jSit3SsDK72s7n7rsSHd+Gs= gorm.io/plugin/opentelemetry v0.1.16 h1:Kypj2YYAliJqkIczDZDde6P6sFMhKSlG5IpngMFQGpc= gorm.io/plugin/opentelemetry v0.1.16/go.mod h1:P3RmTeZXT+9n0F1ccUqR5uuTvEXDxF8k2UpO7mTIB2Y= +mellium.im/sasl v0.3.1 h1:wE0LW6g7U83vhvxjC1IY8DnXM+EU095yeo8XClvCdfo= +mellium.im/sasl v0.3.1/go.mod h1:xm59PUYpZHhgQ9ZqoJ5QaCqzWMi8IeS49dhp6plPCzw= diff --git a/src/server/api/go/internal/bootstrap/container.go b/src/server/api/go/internal/bootstrap/container.go index 3fc51bcad..386e26660 100644 --- a/src/server/api/go/internal/bootstrap/container.go +++ b/src/server/api/go/internal/bootstrap/container.go @@ -3,6 +3,7 @@ package bootstrap import ( "context" "crypto/tls" + "fmt" "strings" "time" @@ -53,6 +54,9 @@ func BuildContainer() *do.Injector { // ALTER TABLE agent_skills DROP COLUMN IF EXISTS asset_meta; // ALTER TABLE agent_skills DROP COLUMN IF EXISTS file_index; if cfg.Database.AutoMigrate { + // Ensure pgvector extension exists + _ = d.Exec("CREATE EXTENSION IF NOT EXISTS vector") + _ = d.AutoMigrate( &model.Project{}, &model.User{}, @@ -66,6 +70,19 @@ func BuildContainer() *do.Injector { &model.AgentSkills{}, &model.SandboxLog{}, ) + + // Create the tasks.embedding vector column with the configured dimension. + dim := cfg.Embedding.TaskVectorDim + _ = d.Exec(fmt.Sprintf( + `DO $$ BEGIN + IF NOT EXISTS ( + SELECT 1 FROM information_schema.columns + WHERE table_name='tasks' AND column_name='embedding' + ) THEN + ALTER TABLE tasks ADD COLUMN embedding vector(%d); + END IF; + END $$;`, dim, + )) } // ensure default project exists diff --git a/src/server/api/go/internal/config/config.go b/src/server/api/go/internal/config/config.go index bd39371e5..ba7a1cb14 100644 --- a/src/server/api/go/internal/config/config.go +++ b/src/server/api/go/internal/config/config.go @@ -86,6 +86,10 @@ type ArtifactCfg struct { MaxUploadSizeBytes int64 // Maximum file upload size in bytes } +type EmbeddingCfg struct { + TaskVectorDim int +} + type Config struct { App AppCfg Root RootCfg @@ -97,6 +101,7 @@ type Config struct { Core CoreCfg Telemetry TelemetryCfg Artifact ArtifactCfg + Embedding EmbeddingCfg } func setDefaults(v *viper.Viper) { @@ -127,6 +132,7 @@ func setDefaults(v *viper.Viper) { v.SetDefault("telemetry.enabled", true) v.SetDefault("telemetry.sampleRatio", 1.0) // Default 100% sampling v.SetDefault("artifact.maxUploadSizeBytes", 16777216) // Default 16MB (16 * 1024 * 1024 bytes) + v.SetDefault("embedding.taskVectorDim", 1536) } func Load() (*Config, error) { diff --git a/src/server/api/go/internal/infra/httpclient/core.go b/src/server/api/go/internal/infra/httpclient/core.go index f48d83124..c6e70bc6a 100644 --- a/src/server/api/go/internal/infra/httpclient/core.go +++ b/src/server/api/go/internal/infra/httpclient/core.go @@ -394,3 +394,51 @@ func (c *CoreClient) UploadSandboxFile(ctx context.Context, projectID, sandboxID return &result, nil } + +// SessionSearchResponse represents the response from session search +type SessionSearchResponse struct { + SessionIDs []uuid.UUID `json:"session_ids"` +} + +// SessionSearch calls the session search endpoint in Python Core +func (c *CoreClient) SessionSearch(ctx context.Context, userID string, projectID string, query string, limit int) (*SessionSearchResponse, error) { + endpoint := fmt.Sprintf("%s/api/v1/sessions/search", c.BaseURL) + + httpReq, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil) + if err != nil { + return nil, fmt.Errorf("create request: %w", err) + } + + // Add query parameters + q := httpReq.URL.Query() + q.Add("user_id", userID) + q.Add("project_id", projectID) + q.Add("query", query) + q.Add("limit", fmt.Sprintf("%d", limit)) + httpReq.URL.RawQuery = q.Encode() + + resp, err := c.HTTPClient.Do(httpReq) + if err != nil { + return nil, fmt.Errorf("do request: %w", err) + } + defer resp.Body.Close() + + respBody, err := io.ReadAll(resp.Body) + if err != nil { + return nil, fmt.Errorf("read response body: %w", err) + } + + if resp.StatusCode != http.StatusOK { + c.Logger.Error("session_search request failed", + zap.Int("status_code", resp.StatusCode), + zap.String("body", string(respBody))) + return nil, fmt.Errorf("request failed with status %d: %s", resp.StatusCode, string(respBody)) + } + + var result SessionSearchResponse + if err := sonic.Unmarshal(respBody, &result); err != nil { + return nil, fmt.Errorf("unmarshal response: %w", err) + } + + return &result, nil +} diff --git a/src/server/api/go/internal/modules/handler/session.go b/src/server/api/go/internal/modules/handler/session.go index f7718701c..269c112e9 100644 --- a/src/server/api/go/internal/modules/handler/session.go +++ b/src/server/api/go/internal/modules/handler/session.go @@ -778,3 +778,45 @@ func (h *SessionHandler) PatchConfigs(c *gin.Context) { c.JSON(http.StatusOK, serializer.Response{Data: PatchSessionConfigsResp{Configs: updatedConfigs}}) } + +type SessionSearchReq struct { + Query string `form:"query" json:"query" binding:"required" example:"find conversations about authentication"` + UserID string `form:"user_id" json:"user_id" binding:"required" example:"123e4567-e89b-12d3-a456-426614174000"` + ProjectID string `form:"project_id" json:"project_id" binding:"required" example:"123e4567-e89b-12d3-a456-426614174001"` + Limit int `form:"limit,default=10" json:"limit" binding:"omitempty,min=1,max=100" example:"10"` +} + +// SessionSearch godoc +// +// @Summary Search sessions +// @Description Search for sessions by semantic similarity to a query string +// @Tags session +// @Accept json +// @Produce json +// @Param query query string true "Search query text" +// @Param user_id query string true "User ID" +// @Param project_id query string true "Project ID" +// @Param limit query integer false "Maximum number of results (1-100, default 10)" +// @Security BearerAuth +// @Success 200 {object} serializer.Response{data=httpclient.SessionSearchResponse} +// @Router /sessions/search [get] +func (h *SessionHandler) SessionSearch(c *gin.Context) { + req := SessionSearchReq{} + if err := c.ShouldBind(&req); err != nil { + c.JSON(http.StatusBadRequest, serializer.ParamErr("", err)) + return + } + + limit := req.Limit + if limit == 0 { + limit = 10 + } + + result, err := h.coreClient.SessionSearch(c.Request.Context(), req.UserID, req.ProjectID, req.Query, limit) + if err != nil { + c.JSON(http.StatusInternalServerError, serializer.Err(http.StatusInternalServerError, "failed to search sessions", err)) + return + } + + c.JSON(http.StatusOK, serializer.Response{Data: result}) +} diff --git a/src/server/api/go/internal/modules/model/message.go b/src/server/api/go/internal/modules/model/message.go index 72fee627a..60128f9e1 100644 --- a/src/server/api/go/internal/modules/model/message.go +++ b/src/server/api/go/internal/modules/model/message.go @@ -189,18 +189,17 @@ type Message struct { PartsAssetMeta datatypes.JSONType[Asset] `gorm:"type:jsonb;not null" swaggertype:"-" json:"-"` Parts []Part `gorm:"-" swaggertype:"array,object" json:"parts"` - TaskID *uuid.UUID `gorm:"type:uuid;index" json:"task_id"` + // Message <-> Task + TaskID *uuid.UUID `gorm:"type:uuid;index:ix_message_task_id;constraint:OnDelete:SET NULL,OnUpdate:CASCADE;" json:"task_id"` + Task *Task `gorm:"foreignKey:TaskID;references:ID" json:"-"` - SessionTaskProcessStatus string `gorm:"type:text;not null;default:'pending';check:session_task_process_status IN ('success','failed','running','pending')" json:"session_task_process_status"` + SessionTaskProcessStatus string `gorm:"type:text;not null;default:'pending';check:session_task_process_status IN ('pending', 'running', 'success', 'failed')" json:"session_task_process_status"` CreatedAt time.Time `gorm:"autoCreateTime;not null;default:CURRENT_TIMESTAMP;index:idx_session_created,priority:2,sort:desc" json:"created_at"` UpdatedAt time.Time `gorm:"autoUpdateTime;not null;default:CURRENT_TIMESTAMP" json:"updated_at"` // Message <-> Session Session *Session `gorm:"foreignKey:SessionID;references:ID;constraint:OnDelete:CASCADE,OnUpdate:CASCADE;" json:"-"` - - // Message <-> Task - Task *Task `gorm:"foreignKey:TaskID;references:ID;constraint:OnDelete:SET NULL,OnUpdate:CASCADE;" json:"-"` } func (Message) TableName() string { return "messages" } diff --git a/src/server/api/go/internal/modules/model/task.go b/src/server/api/go/internal/modules/model/task.go index 8ac1554ff..9b9c5e6c3 100644 --- a/src/server/api/go/internal/modules/model/task.go +++ b/src/server/api/go/internal/modules/model/task.go @@ -7,6 +7,7 @@ import ( "time" "github.com/google/uuid" + "github.com/pgvector/pgvector-go" ) type Task struct { @@ -19,6 +20,9 @@ type Task struct { Status string `gorm:"type:text;not null;default:'pending';check:status IN ('success','failed','running','pending');index:ix_task_session_id_status,priority:2" json:"status"` IsPlanning bool `gorm:"not null;default:false" json:"is_planning"` + // Embedding vector for semantic search + Embedding pgvector.Vector `gorm:"type:vector;-:migration" json:"-"` + CreatedAt time.Time `gorm:"autoCreateTime;not null;default:CURRENT_TIMESTAMP" json:"created_at"` UpdatedAt time.Time `gorm:"autoUpdateTime;not null;default:CURRENT_TIMESTAMP" json:"updated_at"` diff --git a/src/server/api/go/internal/modules/repo/session_test.go b/src/server/api/go/internal/modules/repo/session_test.go index 428445538..1fb6db608 100644 --- a/src/server/api/go/internal/modules/repo/session_test.go +++ b/src/server/api/go/internal/modules/repo/session_test.go @@ -25,6 +25,9 @@ func setupSessionTestDB(t *testing.T) *gorm.DB { return nil } + // Ensure pgvector extension exists + db.Exec("CREATE EXTENSION IF NOT EXISTS vector") + // Auto migrate all required tables err = db.AutoMigrate( &model.Project{}, diff --git a/src/server/api/go/internal/router/router.go b/src/server/api/go/internal/router/router.go index 6ef84f62c..dd90640d7 100644 --- a/src/server/api/go/internal/router/router.go +++ b/src/server/api/go/internal/router/router.go @@ -61,6 +61,12 @@ func NewRouter(d RouterDeps) *gin.Engine { // ping endpoint v1.GET("/ping", func(c *gin.Context) { c.JSON(http.StatusOK, serializer.Response{Msg: "pong"}) }) + // Sessions search (project-level, without session_id) + sessions := v1.Group("/sessions") + { + sessions.GET("/search", d.SessionHandler.SessionSearch) + } + session := v1.Group("/session") { session.GET("", d.SessionHandler.GetSessions) diff --git a/src/server/core/acontext_core/infra/db.py b/src/server/core/acontext_core/infra/db.py index 716a5ff89..aad349a4c 100644 --- a/src/server/core/acontext_core/infra/db.py +++ b/src/server/core/acontext_core/infra/db.py @@ -238,6 +238,39 @@ async def init_database() -> None: assert await DB_CLIENT.health_check(), "Database health check failed" logger.info(f"Database created successfully {DB_CLIENT.get_pool_status()}") + # Validate vector dimension matches configuration + expected_dim = DEFAULT_CORE_CONFIG.task_embedding_dim + + async with DB_CLIENT.get_session_context() as session: + result = await session.execute(text("SELECT 1 FROM pg_extension WHERE extname = 'vector'")) + if not result.scalar(): + logger.warning("pgvector extension not installed. Skipping dimension check.") + return + + try: + # Check the defined dimension of the 'embedding' column in 'tasks' + sql = text(""" + SELECT atttypmod + FROM pg_attribute + WHERE attrelid = 'tasks'::regclass + AND attname = 'embedding' + """) + result = await session.execute(sql) + val = result.scalar() + + if val is not None: + if val != expected_dim: + logger.warning( + f"Database 'tasks.embedding' dimension ({val}) does not match " + f"config ({expected_dim})! This may cause runtime errors." + ) + else: + logger.info("Database vector dimension matches config.") + else: + logger.info("Tasks table or embedding column not found. Skipping dimension check (fresh install?).") + except Exception as e: + logger.warning(f"Could not verify database vector dimension: {e}") + async def close_database() -> None: """Close database connections.""" diff --git a/src/server/core/acontext_core/schema/config.py b/src/server/core/acontext_core/schema/config.py index 2e5b6aa64..d1614251c 100644 --- a/src/server/core/acontext_core/schema/config.py +++ b/src/server/core/acontext_core/schema/config.py @@ -24,6 +24,10 @@ class CoreConfig(BaseModel): llm_simple_model: str = "gpt-4.1" + # Embedding Configuration + task_embedding_model: str = "text-embedding-3-small" + task_embedding_dim: int = 1536 + # Core Configuration logging_format: str = "text" logging_level: str = "INFO" diff --git a/src/server/core/acontext_core/schema/orm/task.py b/src/server/core/acontext_core/schema/orm/task.py index ee0a33977..38e5ea34e 100644 --- a/src/server/core/acontext_core/schema/orm/task.py +++ b/src/server/core/acontext_core/schema/orm/task.py @@ -11,9 +11,12 @@ ) from sqlalchemy.orm import relationship from sqlalchemy.dialects.postgresql import JSONB, UUID -from typing import TYPE_CHECKING, List +from typing import TYPE_CHECKING, List, Optional + from .base import ORM_BASE, CommonMixin from ..utils import asUUID +from pgvector.sqlalchemy import Vector +from ...env import DEFAULT_CORE_CONFIG if TYPE_CHECKING: from .project import Project @@ -48,42 +51,42 @@ class Task(CommonMixin): metadata={ "db": Column( UUID(as_uuid=True), - ForeignKey("sessions.id", ondelete="CASCADE"), + ForeignKey("sessions.id", ondelete="CASCADE", onupdate="CASCADE"), nullable=False, ) - } + }, ) project_id: asUUID = field( metadata={ "db": Column( UUID(as_uuid=True), - ForeignKey("projects.id", ondelete="CASCADE"), + ForeignKey("projects.id", ondelete="CASCADE", onupdate="CASCADE"), nullable=False, ) - } + }, ) order: int = field(metadata={"db": Column(Integer, nullable=False)}) - data: dict = field(metadata={"db": Column(JSONB, nullable=False)}) + data: dict = field( + default_factory=dict, metadata={"db": Column(JSONB, nullable=False)} + ) status: str = field( default="pending", - metadata={ - "db": Column( - String, - nullable=False, - server_default="pending", - ) - }, + metadata={"db": Column(String, nullable=False, server_default="pending")}, ) is_planning: bool = field( default=False, - metadata={ - "db": Column(Boolean, nullable=False, default=False, server_default="false") - }, + metadata={"db": Column(Boolean, nullable=False, server_default="false")}, + ) + + # Embedding vector for semantic search + embedding: Optional[List[float]] = field( + default=None, + metadata={"db": Column(Vector(DEFAULT_CORE_CONFIG.task_embedding_dim), nullable=True)}, ) # Relationships diff --git a/src/server/core/acontext_core/service/data/session_search_service.py b/src/server/core/acontext_core/service/data/session_search_service.py new file mode 100644 index 000000000..ddd9912d1 --- /dev/null +++ b/src/server/core/acontext_core/service/data/session_search_service.py @@ -0,0 +1,72 @@ +from typing import List +from sqlalchemy import text +from sqlalchemy.ext.asyncio import AsyncSession + +from ...schema.result import Result +from ...schema.utils import asUUID +from ...llm.embeddings.openai_embedding import openai_embedding +from ...env import LOG, DEFAULT_CORE_CONFIG + +async def search_sessions_by_task_query( + db_session: AsyncSession, + user_id: asUUID, + project_id: asUUID, + query: str, + topk: int = 10, + threshold: float = 0.8, +) -> Result[List[asUUID]]: + """ + Search for sessions by semantic similarity of their Tasks to a query string. + + Args: + db_session: Database session + user_id: User ID to scope the search + project_id: Project ID to scope the search + query: Search query text + topk: Maximum number of results to return + threshold: Cosine distance threshold (lower = more similar) + + Returns: + Result containing list of unique session_ids ordered by relevance + """ + try: + embedding_result = await openai_embedding( + model=DEFAULT_CORE_CONFIG.task_embedding_model, + texts=[query], + phase="query", + ) + query_embedding = embedding_result.embedding[0] + + # Perform vector similarity search using pgvector on TASKS table + sql = text(""" + SELECT DISTINCT t.session_id, MIN(t.embedding <=> :query_embedding::vector) as distance + FROM tasks t + JOIN sessions s ON t.session_id = s.id + WHERE s.user_id = :user_id + AND s.project_id = :project_id + AND t.embedding IS NOT NULL + AND (t.embedding <=> :query_embedding::vector) < :threshold + GROUP BY t.session_id + ORDER BY distance ASC + LIMIT :topk + """) + + result = await db_session.execute( + sql, + { + "query_embedding": str(query_embedding), + "user_id": str(user_id), + "project_id": str(project_id), + "threshold": threshold, + "topk": topk, + }, + ) + + session_ids = [row[0] for row in result.fetchall()] + + LOG.info(f"Session search (via Tasks) found {len(session_ids)} results for query: {query[:50]}...") + return Result.resolve(session_ids) + + except Exception as e: + LOG.error(f"Error searching sessions by task: {e}") + return Result.reject(f"Error searching sessions: {e}") diff --git a/src/server/core/api.py b/src/server/core/api.py index dd5a140c1..e51a0bace 100644 --- a/src/server/core/api.py +++ b/src/server/core/api.py @@ -12,7 +12,7 @@ shutdown_otel_tracing, ) from acontext_core.telemetry.config import TelemetryConfig -from routers import session_router, sandbox_router +from routers import session_router, search_router, sandbox_router # Filter to exclude /health endpoint from uvicorn access logs @@ -85,6 +85,7 @@ async def lifespan(app: FastAPI): # Include routers app.include_router(session_router) +app.include_router(search_router) app.include_router(sandbox_router) # Instrument FastAPI app after creation and route registration diff --git a/src/server/core/routers/__init__.py b/src/server/core/routers/__init__.py index 9553dc8d0..f19001517 100644 --- a/src/server/core/routers/__init__.py +++ b/src/server/core/routers/__init__.py @@ -1,7 +1,8 @@ -from .session import router as session_router +from .session import router as session_router, search_router from .sandbox import router as sandbox_router __all__ = [ "session_router", + "search_router", "sandbox_router", ] diff --git a/src/server/core/routers/session.py b/src/server/core/routers/session.py index 27d216db1..73bb6763f 100644 --- a/src/server/core/routers/session.py +++ b/src/server/core/routers/session.py @@ -1,8 +1,14 @@ -from fastapi import APIRouter, Path +from typing import List +from fastapi import APIRouter, Path, Query +from fastapi.exceptions import HTTPException +from pydantic import BaseModel + from acontext_core.env import LOG from acontext_core.schema.api.response import Flag from acontext_core.schema.utils import asUUID from acontext_core.service.session_message import flush_session_message_blocking +from acontext_core.service.data.session_search_service import search_sessions_by_task_query +from acontext_core.infra.db import DB_CLIENT router = APIRouter(prefix="/api/v1/project/{project_id}/session/{session_id}", tags=["session"]) @@ -18,3 +24,39 @@ async def session_flush( LOG.info(f"Flushing session {session_id} for project {project_id}") r = await flush_session_message_blocking(project_id, session_id) return Flag(status=r.error.status.value, errmsg=r.error.errmsg) + + +# Search router for project-level session search +class SessionSearchResponse(BaseModel): + session_ids: List[str] + + +search_router = APIRouter(prefix="/api/v1/sessions", tags=["session_search"]) + + +@search_router.get("/search") +async def session_search( + user_id: asUUID = Query(..., description="User ID to search within"), + project_id: asUUID = Query(..., description="Project ID to search within"), + query: str = Query(..., description="Search query text"), + limit: int = Query(10, ge=1, le=100, description="Maximum results to return"), +) -> SessionSearchResponse: + """ + Uses vector embeddings on Tasks to find sessions with relevant context. + """ + LOG.info(f"Searching sessions in project {project_id} for user {user_id} with query: {query[:50]}...") + + async with DB_CLIENT.get_session_context() as db_session: + result = await search_sessions_by_task_query( + db_session, + user_id, + project_id, + query, + topk=limit, + ) + + if not result.ok(): + raise HTTPException(status_code=500, detail=result.error.errmsg) + + session_ids = [str(sid) for sid in result.data] + return SessionSearchResponse(session_ids=session_ids)