correction des pb de space : le spaceId est maintenant dans l'url uniquement
This commit is contained in:
1 parent
c52460bac0
commit
0bb8d1cb46
41 files changed
+1463
-121
No files matched your search
@@ -13,6 +13,7 @@ type Database interface {
|
||||
UserRepository() domain.UserRepository
|
||||
FileRepository() domain.FileRepository
|
||||
ProviderRepository() domain.ProviderRepository
|
||||
AgentRepository() domain.AgentRepository
|
||||
Migrate() error
|
||||
Close()
|
||||
}
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
package orm
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func buildInsertSql(tablename string, fields []dbfield, dialect Dialect) string {
|
||||
cols := make([]string, len(fields))
|
||||
placeholders := make([]string, len(fields))
|
||||
for i, f := range fields {
|
||||
cols[i] = f.dbname
|
||||
placeholders[i] = fmt.Sprintf("$%d", i+1)
|
||||
}
|
||||
return fmt.Sprintf(
|
||||
"insert into %s (%s) values (%s)",
|
||||
tablename,
|
||||
strings.Join(cols, ","),
|
||||
strings.Join(placeholders, ","),
|
||||
)
|
||||
}
|
||||
|
||||
func buildUpdateSql(tablename string, fields []dbfield, dialect Dialect) string {
|
||||
predicates := make([]string, len(fields)-1)
|
||||
var where string
|
||||
i := 0
|
||||
for _, f := range fields {
|
||||
if f.dbname == "id" {
|
||||
where = fmt.Sprintf("id=$%d and _version=$%d", len(fields), len(fields)+1)
|
||||
} else {
|
||||
predicates[i] = fmt.Sprintf("%s=$%d", f.dbname, i+1)
|
||||
i++
|
||||
}
|
||||
}
|
||||
return fmt.Sprintf(
|
||||
"update %s set %s where %s",
|
||||
tablename,
|
||||
strings.Join(predicates, ","),
|
||||
where,
|
||||
)
|
||||
}
|
||||
|
||||
func buildSelectbyidSql(tablename string, fields []dbfield, dialect Dialect) string {
|
||||
columns := make([]string, len(fields))
|
||||
var where string
|
||||
for i, f := range fields {
|
||||
columns[i] = fmt.Sprintf("%s", f.dbname)
|
||||
}
|
||||
where = "id=$1"
|
||||
return fmt.Sprintf(
|
||||
"select %s from %s where %s",
|
||||
strings.Join(columns, ","),
|
||||
tablename,
|
||||
where,
|
||||
)
|
||||
}
|
||||
|
||||
func buildDeletebyidSql(tablename string, fields []dbfield, dialect Dialect) string {
|
||||
where := "id=$1"
|
||||
return fmt.Sprintf(
|
||||
"delete from %s where %s",
|
||||
tablename,
|
||||
where,
|
||||
)
|
||||
}
|
||||
|
||||
func buildSelectWhere(tablename string, fields []dbfield, dialect Dialect) string {
|
||||
columns := make([]string, len(fields))
|
||||
for i, f := range fields {
|
||||
columns[i] = fmt.Sprintf("%s", f.dbname)
|
||||
}
|
||||
return fmt.Sprintf(
|
||||
"select %s from %s",
|
||||
strings.Join(columns, ","),
|
||||
tablename,
|
||||
)
|
||||
}
|
||||
|
||||
func buildDeleteWhere(tablename string, fields []dbfield, dialect Dialect) string {
|
||||
return fmt.Sprintf(
|
||||
"delete from %s",
|
||||
tablename,
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,289 @@
|
||||
package orm
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type SqliteDialect struct{}
|
||||
|
||||
func NewSqliteDialect() *SqliteDialect {
|
||||
return &SqliteDialect{}
|
||||
}
|
||||
|
||||
func convertType(f dbfield) string {
|
||||
switch f.stype {
|
||||
case "string", "*string":
|
||||
return "text"
|
||||
case "bool", "int", "*int", "time.Time", "*time.Time":
|
||||
return "numeric"
|
||||
case "[]byte":
|
||||
return "blob"
|
||||
}
|
||||
return "text"
|
||||
}
|
||||
|
||||
func defaultValue(f dbfield) string {
|
||||
if f.dbname == "_date_created" {
|
||||
return "current_timestamp"
|
||||
}
|
||||
switch f.stype {
|
||||
case "string":
|
||||
return "''"
|
||||
case "bool", "int", "time.Time":
|
||||
return "0"
|
||||
}
|
||||
return "''"
|
||||
}
|
||||
|
||||
func (d *SqliteDialect) TableExists(db *sql.DB, tableName string) (bool, error) {
|
||||
sql := "select name from sqlite_schema where type='table' and name=$1"
|
||||
rows, err := db.Query(sql, tableName)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if rows.Next() {
|
||||
rows.Close()
|
||||
return true, nil
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
func (d *SqliteDialect) ColumnSpec(f dbfield) (string, error) {
|
||||
if f.dbname == "id" {
|
||||
return "id text not null primary key", nil
|
||||
}
|
||||
if strings.HasPrefix(f.stype, "*") || strings.HasPrefix(f.stype, "[]") {
|
||||
return fmt.Sprintf("%s %s", f.dbname, convertType(f)), nil
|
||||
}
|
||||
return fmt.Sprintf("%s %s not null default %s", f.dbname, convertType(f), defaultValue(f)), nil
|
||||
}
|
||||
|
||||
// ScanResult scans a single row into the provided fields slice.
|
||||
// For SQLite, this handles time.Time fields which are returned as strings.
|
||||
func (d *SqliteDialect) ScanResult(row interface{ Scan(dest ...any) error }, fields []any, fieldTypes []dbfield) error {
|
||||
// SQLite returns time.Time as strings, so we need special handling
|
||||
// Check if any field is a time type
|
||||
hasTimeField := false
|
||||
for _, f := range fieldTypes {
|
||||
if f.stype == "time.Time" || f.stype == "*time.Time" {
|
||||
hasTimeField = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !hasTimeField {
|
||||
// No time fields, use standard scan
|
||||
return row.Scan(fields...)
|
||||
}
|
||||
|
||||
// We have time fields - need custom handling
|
||||
// Create temp destinations for all fields
|
||||
tempFields := make([]any, len(fields))
|
||||
fieldInfos := make([]struct {
|
||||
index int
|
||||
dest any
|
||||
isTime bool
|
||||
isTimePtr bool
|
||||
}, len(fields))
|
||||
|
||||
for i, f := range fieldTypes {
|
||||
isTime := f.stype == "time.Time"
|
||||
isTimePtr := f.stype == "*time.Time"
|
||||
|
||||
fieldInfos[i] = struct {
|
||||
index int
|
||||
dest any
|
||||
isTime bool
|
||||
isTimePtr bool
|
||||
}{
|
||||
index: i,
|
||||
dest: fields[i],
|
||||
isTime: isTime,
|
||||
isTimePtr: isTimePtr,
|
||||
}
|
||||
|
||||
if isTime || isTimePtr {
|
||||
var s sql.NullString
|
||||
tempFields[i] = &s
|
||||
} else {
|
||||
tempFields[i] = fields[i]
|
||||
}
|
||||
}
|
||||
|
||||
if err := row.Scan(tempFields...); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Copy time values back to original destinations
|
||||
for _, info := range fieldInfos {
|
||||
if info.isTime || info.isTimePtr {
|
||||
s := tempFields[info.index].(*sql.NullString)
|
||||
if s.Valid {
|
||||
// Try to parse the string as time
|
||||
t, err := tryParseTime(s.String)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to parse time field: %w", err)
|
||||
}
|
||||
|
||||
// Set the value in the original destination
|
||||
destVal := reflect.ValueOf(info.dest).Elem()
|
||||
if info.isTime {
|
||||
// dest is *time.Time, set the value
|
||||
destVal.Set(reflect.ValueOf(t))
|
||||
} else if info.isTimePtr {
|
||||
// dest is **time.Time, allocate and set
|
||||
ptr := reflect.New(destVal.Type().Elem())
|
||||
ptr.Elem().Set(reflect.ValueOf(t))
|
||||
destVal.Set(ptr)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ScanRows scans multiple rows using the dialect-specific logic
|
||||
func (d *SqliteDialect) ScanRows(rows *sql.Rows, fieldDest func() []any, fieldTypes []dbfield) ([]any, error) {
|
||||
// Check if any field is a time type
|
||||
hasTimeField := false
|
||||
for _, f := range fieldTypes {
|
||||
if f.stype == "time.Time" || f.stype == "*time.Time" {
|
||||
hasTimeField = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !hasTimeField {
|
||||
// No time fields, use standard scan
|
||||
var results []any
|
||||
for rows.Next() {
|
||||
fields := fieldDest()
|
||||
if err := rows.Scan(fields...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
results = append(results, fields...)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return results, nil
|
||||
}
|
||||
|
||||
// We have time fields - need custom handling
|
||||
var results []any
|
||||
for rows.Next() {
|
||||
fields := fieldDest()
|
||||
|
||||
// Create temp destinations for all fields
|
||||
tempFields := make([]any, len(fields))
|
||||
fieldInfos := make([]struct {
|
||||
index int
|
||||
dest any
|
||||
isTime bool
|
||||
isTimePtr bool
|
||||
}, len(fields))
|
||||
|
||||
for i, f := range fieldTypes {
|
||||
isTime := f.stype == "time.Time"
|
||||
isTimePtr := f.stype == "*time.Time"
|
||||
|
||||
fieldInfos[i] = struct {
|
||||
index int
|
||||
dest any
|
||||
isTime bool
|
||||
isTimePtr bool
|
||||
}{
|
||||
index: i,
|
||||
dest: fields[i],
|
||||
isTime: isTime,
|
||||
isTimePtr: isTimePtr,
|
||||
}
|
||||
|
||||
if isTime || isTimePtr {
|
||||
var s sql.NullString
|
||||
tempFields[i] = &s
|
||||
} else {
|
||||
tempFields[i] = fields[i]
|
||||
}
|
||||
}
|
||||
|
||||
if err := rows.Scan(tempFields...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Copy time values back to original destinations
|
||||
for _, info := range fieldInfos {
|
||||
if info.isTime || info.isTimePtr {
|
||||
s := tempFields[info.index].(*sql.NullString)
|
||||
if s.Valid {
|
||||
t, err := tryParseTime(s.String)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to parse time field: %w", err)
|
||||
}
|
||||
|
||||
// Set the value in the original destination
|
||||
destVal := reflect.ValueOf(info.dest).Elem()
|
||||
if info.isTime {
|
||||
// dest is *time.Time, set the value
|
||||
destVal.Set(reflect.ValueOf(t))
|
||||
} else if info.isTimePtr {
|
||||
// dest is **time.Time, allocate and set
|
||||
ptr := reflect.New(destVal.Type().Elem())
|
||||
ptr.Elem().Set(reflect.ValueOf(t))
|
||||
destVal.Set(ptr)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
results = append(results, fields...)
|
||||
}
|
||||
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return results, nil
|
||||
}
|
||||
|
||||
// tryParseTime attempts to parse a time string in various common formats
|
||||
// SQLite driver returns time as "2006-01-02 15:04:05.999999999 -0700 MST m=+0.000000000"
|
||||
// We need to strip the monotonic clock part (m=...) before parsing
|
||||
func tryParseTime(s string) (time.Time, error) {
|
||||
// Remove monotonic clock part if present
|
||||
// Format: "2006-01-02 15:04:05.999999999 -0700 MST m=+0.000000000"
|
||||
if idx := strings.Index(s, " m="); idx != -1 {
|
||||
s = s[:idx]
|
||||
}
|
||||
|
||||
// Try common formats
|
||||
formats := []string{
|
||||
time.RFC3339,
|
||||
time.RFC3339Nano,
|
||||
"2006-01-02T15:04:05Z07:00",
|
||||
"2006-01-02 15:04:05.999999999-07:00",
|
||||
"2006-01-02 15:04:05",
|
||||
"2006-01-02T15:04:05",
|
||||
"2006-01-02 15:04:05+00:00",
|
||||
// Go's default time.String() format (without monotonic clock)
|
||||
"2006-01-02 15:04:05.999999999 -0700 MST",
|
||||
// SQLite without timezone
|
||||
"2006-01-02 15:04:05.999999",
|
||||
}
|
||||
|
||||
var err error
|
||||
var t time.Time
|
||||
for _, format := range formats {
|
||||
t, err = time.Parse(format, s)
|
||||
if err == nil {
|
||||
return t, nil
|
||||
}
|
||||
}
|
||||
|
||||
return time.Time{}, fmt.Errorf("unable to parse time: %s", s)
|
||||
}
|
||||
@@ -0,0 +1,373 @@
|
||||
package orm
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type dbfield struct {
|
||||
name string
|
||||
stype string
|
||||
dbname string
|
||||
}
|
||||
|
||||
type Dialect interface {
|
||||
TableExists(db *sql.DB, tableName string) (bool, error)
|
||||
ColumnSpec(f dbfield) (string, error)
|
||||
ScanResult(row interface{ Scan(dest ...any) error }, fields []any, fieldTypes []dbfield) error
|
||||
}
|
||||
|
||||
type Table[T any] struct {
|
||||
db *sql.DB
|
||||
dialect Dialect
|
||||
tablename string
|
||||
fields []dbfield
|
||||
insertsql string
|
||||
updatesql string
|
||||
selectbyid string
|
||||
deletebyid string
|
||||
selectwhere string
|
||||
deletewhere string
|
||||
}
|
||||
|
||||
func NewTable[T any](db *sql.DB, dialect Dialect, tablename string) (Table[T], error) {
|
||||
|
||||
var instance T
|
||||
|
||||
t := reflect.TypeOf(instance)
|
||||
|
||||
fields := make([]dbfield, 0)
|
||||
for field := range t.Fields() {
|
||||
f := dbfield{
|
||||
name: field.Name,
|
||||
stype: field.Type.String(),
|
||||
dbname: field.Tag.Get("db"),
|
||||
}
|
||||
fields = append(fields, f)
|
||||
}
|
||||
|
||||
return Table[T]{
|
||||
db: db,
|
||||
dialect: dialect,
|
||||
tablename: tablename,
|
||||
fields: fields,
|
||||
insertsql: buildInsertSql(tablename, fields, dialect),
|
||||
updatesql: buildUpdateSql(tablename, fields, dialect),
|
||||
selectbyid: buildSelectbyidSql(tablename, fields, dialect),
|
||||
deletebyid: buildDeletebyidSql(tablename, fields, dialect),
|
||||
selectwhere: buildSelectWhere(tablename, fields, dialect),
|
||||
deletewhere: buildDeleteWhere(tablename, fields, dialect),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (t Table[T]) Sync() error {
|
||||
if exists, err := t.dialect.TableExists(t.db, t.tablename); exists {
|
||||
return nil
|
||||
} else if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
columns := make([]string, len(t.fields))
|
||||
for i, f := range t.fields {
|
||||
c, err := t.dialect.ColumnSpec(f)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
columns[i] = c
|
||||
}
|
||||
sql := fmt.Sprintf("create table %s (%s)", t.tablename, strings.Join(columns, ","))
|
||||
_, err := t.db.Exec(sql)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t Table[T]) Get(id string) (*T, error) {
|
||||
row := t.db.QueryRow(t.selectbyid, id)
|
||||
if row == nil {
|
||||
return nil, sql.ErrNoRows
|
||||
}
|
||||
var obj T
|
||||
|
||||
// Create field destinations
|
||||
ptrVal := reflect.ValueOf(&obj).Elem()
|
||||
fields := make([]any, len(t.fields))
|
||||
for i, f := range t.fields {
|
||||
fieldValue := ptrVal.FieldByName(f.name)
|
||||
if !fieldValue.IsValid() {
|
||||
return nil, fmt.Errorf("field %s not found in struct", f.name)
|
||||
}
|
||||
fields[i] = fieldValue.Addr().Interface()
|
||||
}
|
||||
|
||||
// Use dialect to scan with proper type handling
|
||||
if err := t.dialect.ScanResult(row, fields, t.fields); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, sql.ErrNoRows
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &obj, nil
|
||||
}
|
||||
|
||||
func (t Table[T]) SelectOne(where string, args ...any) (*T, error) {
|
||||
query := t.selectwhere + " where " + where + " limit 1"
|
||||
row := t.db.QueryRow(query, args...)
|
||||
if row == nil {
|
||||
return nil, sql.ErrNoRows
|
||||
}
|
||||
|
||||
var obj T
|
||||
ptrVal := reflect.ValueOf(&obj).Elem()
|
||||
fields := make([]any, len(t.fields))
|
||||
for i, f := range t.fields {
|
||||
fieldValue := ptrVal.FieldByName(f.name)
|
||||
if !fieldValue.IsValid() {
|
||||
return nil, fmt.Errorf("field %s not found in struct", f.name)
|
||||
}
|
||||
fields[i] = fieldValue.Addr().Interface()
|
||||
}
|
||||
|
||||
if err := t.dialect.ScanResult(row, fields, t.fields); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, sql.ErrNoRows
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &obj, nil
|
||||
}
|
||||
|
||||
type query struct {
|
||||
where string
|
||||
args []any
|
||||
orderby string
|
||||
pagenum int
|
||||
pagesize int
|
||||
}
|
||||
|
||||
type Option func(*query)
|
||||
|
||||
func WithWhere(where string, args ...any) Option {
|
||||
return func(q *query) {
|
||||
q.where = where
|
||||
q.args = args
|
||||
}
|
||||
}
|
||||
|
||||
func WithOrder(orderby string) Option {
|
||||
return func(q *query) {
|
||||
q.orderby = orderby
|
||||
}
|
||||
}
|
||||
|
||||
func WithPagination(pagenum, pagesize int) Option {
|
||||
return func(q *query) {
|
||||
q.pagenum = pagenum
|
||||
q.pagesize = pagesize
|
||||
}
|
||||
}
|
||||
|
||||
func (t Table[T]) Select(options ...Option) ([]*T, error) {
|
||||
q := &query{}
|
||||
for _, o := range options {
|
||||
o(q)
|
||||
}
|
||||
query := t.selectwhere
|
||||
if q.where != "" {
|
||||
query += " where " + q.where
|
||||
}
|
||||
if q.orderby != "" {
|
||||
query += " order by " + q.orderby
|
||||
}
|
||||
if q.pagesize != 0 {
|
||||
query += fmt.Sprintf(" limit %d", q.pagesize)
|
||||
if q.pagenum != 0 {
|
||||
query += fmt.Sprintf(" offset %d", q.pagesize*q.pagenum)
|
||||
}
|
||||
}
|
||||
|
||||
var err error
|
||||
var rows *sql.Rows
|
||||
if q.args == nil {
|
||||
rows, err = t.db.Query(query)
|
||||
} else {
|
||||
rows, err = t.db.Query(query, q.args...)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var results []*T
|
||||
|
||||
for rows.Next() {
|
||||
var obj T
|
||||
ptrVal := reflect.ValueOf(&obj).Elem()
|
||||
fields := make([]any, len(t.fields))
|
||||
for i, f := range t.fields {
|
||||
fieldValue := ptrVal.FieldByName(f.name)
|
||||
if !fieldValue.IsValid() {
|
||||
return nil, fmt.Errorf("field %s not found in struct", f.name)
|
||||
}
|
||||
fields[i] = fieldValue.Addr().Interface()
|
||||
}
|
||||
|
||||
// Use dialect to scan with proper type handling
|
||||
// *sql.Rows implements Scanner interface
|
||||
if err := t.dialect.ScanResult(rows, fields, t.fields); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
results = append(results, &obj)
|
||||
}
|
||||
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func (t Table[T]) Delete(id string) error {
|
||||
_, err := t.db.Exec(t.deletebyid, id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (t Table[T]) DeleteWhere(where string, args ...any) error {
|
||||
query := t.deletewhere
|
||||
if where != "" {
|
||||
query += " where " + where
|
||||
}
|
||||
_, err := t.db.Exec(query, args...)
|
||||
return err
|
||||
}
|
||||
|
||||
func (t Table[T]) Insert(value *T) (*T, error) {
|
||||
// Get a settable reflect value of the struct
|
||||
// value is *T, &value is **T, .Elem() gives *T (settables), .Elem() gives T (settables)
|
||||
val := reflect.ValueOf(&value).Elem().Elem()
|
||||
|
||||
// Auto-fill ID field if it exists and is empty
|
||||
if fieldID := val.FieldByName("ID"); fieldID.IsValid() && fieldID.CanSet() && fieldID.Kind() == reflect.String {
|
||||
if fieldID.String() == "" {
|
||||
fieldID.SetString(GenID())
|
||||
}
|
||||
}
|
||||
|
||||
// Auto-fill VersionId field if it exists and is empty
|
||||
if fieldVersion := val.FieldByName("VersionId"); fieldVersion.IsValid() && fieldVersion.CanSet() && fieldVersion.Kind() == reflect.String {
|
||||
if fieldVersion.String() == "" {
|
||||
fieldVersion.SetString(GenID())
|
||||
}
|
||||
}
|
||||
|
||||
// Auto-fill DateCreated or DateCreate field if it exists
|
||||
now := time.Now()
|
||||
if fieldCreate := val.FieldByName("DateCreated"); fieldCreate.IsValid() && fieldCreate.CanSet() && fieldCreate.Type() == reflect.TypeOf(now) {
|
||||
fieldCreate.Set(reflect.ValueOf(now))
|
||||
} else if fieldCreate := val.FieldByName("DateCreate"); fieldCreate.IsValid() && fieldCreate.CanSet() && fieldCreate.Type() == reflect.TypeOf(now) {
|
||||
fieldCreate.Set(reflect.ValueOf(now))
|
||||
}
|
||||
|
||||
vals := make([]any, len(t.fields))
|
||||
for i, f := range t.fields {
|
||||
fieldValue := val.FieldByName(f.name)
|
||||
if !fieldValue.IsValid() {
|
||||
return nil, fmt.Errorf("field %s not found in struct", f.name)
|
||||
}
|
||||
vals[i] = fieldValue.Interface()
|
||||
}
|
||||
|
||||
_, err := t.db.Exec(t.insertsql, vals...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return value, nil
|
||||
}
|
||||
|
||||
func (t Table[T]) Update(value *T) (*T, error) {
|
||||
// Get a settable reflect value of the struct
|
||||
val := reflect.ValueOf(&value).Elem().Elem()
|
||||
|
||||
// Save old VersionId for WHERE clause
|
||||
versionField := val.FieldByName("VersionId")
|
||||
if !versionField.IsValid() {
|
||||
return nil, fmt.Errorf("VersionId field not found")
|
||||
}
|
||||
oldVersionId := versionField.String()
|
||||
|
||||
// Generate new VersionId for optimistic locking
|
||||
newVersionId := GenID()
|
||||
versionField.SetString(newVersionId)
|
||||
|
||||
// Update DateUpdated field if it exists
|
||||
if fieldUpdated := val.FieldByName("DateUpdated"); fieldUpdated.IsValid() && fieldUpdated.CanSet() {
|
||||
if fieldUpdated.Kind() == reflect.Ptr {
|
||||
now := time.Now()
|
||||
if fieldUpdated.IsNil() {
|
||||
fieldUpdated.Set(reflect.New(fieldUpdated.Type().Elem()))
|
||||
}
|
||||
fieldUpdated.Elem().Set(reflect.ValueOf(now))
|
||||
} else if fieldUpdated.Type() == reflect.TypeOf(time.Time{}) {
|
||||
fieldUpdated.Set(reflect.ValueOf(time.Now()))
|
||||
}
|
||||
}
|
||||
|
||||
// Build the values slice for the prepared statement
|
||||
// UPDATE SQL expects: SET fields (excluding ID), then WHERE: ID and VersionId (old version)
|
||||
vals := make([]any, 0, len(t.fields))
|
||||
|
||||
// Add all fields except ID for SET clause (includes new VersionId)
|
||||
for _, f := range t.fields {
|
||||
if f.dbname == "id" {
|
||||
continue // Skip ID field for SET clause
|
||||
}
|
||||
fieldValue := val.FieldByName(f.name)
|
||||
if !fieldValue.IsValid() {
|
||||
return nil, fmt.Errorf("field %s not found in struct", f.name)
|
||||
}
|
||||
vals = append(vals, fieldValue.Interface())
|
||||
}
|
||||
|
||||
// Add ID and OLD VersionId at the end for WHERE clause
|
||||
idField := val.FieldByName("ID")
|
||||
|
||||
if !idField.IsValid() {
|
||||
return nil, fmt.Errorf("ID field not found")
|
||||
}
|
||||
|
||||
vals = append(vals, idField.Interface(), oldVersionId)
|
||||
|
||||
result, err := t.db.Exec(t.updatesql, vals...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Check if any row was affected
|
||||
rowsAffected, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if rowsAffected == 0 {
|
||||
return nil, fmt.Errorf("no rows affected: record not found or version mismatch")
|
||||
}
|
||||
|
||||
return value, nil
|
||||
}
|
||||
|
||||
func (t Table[T]) Debug() {
|
||||
fmt.Println("insert: ", t.insertsql)
|
||||
fmt.Println("update: ", t.updatesql)
|
||||
fmt.Println("selectbyid: ", t.selectbyid)
|
||||
fmt.Println("deletebyid: ", t.deletebyid)
|
||||
fmt.Println("selectwhere: ", t.selectwhere)
|
||||
fmt.Println("deletewhere: ", t.deletewhere)
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
package orm
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
|
||||
"github.com/sixafter/nanoid"
|
||||
)
|
||||
|
||||
func GenID() string {
|
||||
id, err := nanoid.New()
|
||||
if err != nil {
|
||||
// Fallback to random hex
|
||||
b := make([]byte, 16)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "fallback-id"
|
||||
}
|
||||
return hex.EncodeToString(b)
|
||||
}
|
||||
return id.String()
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
package turso
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"log"
|
||||
|
||||
"trankilou.fr/lassistanoque/backend/internal/adapter/database/orm"
|
||||
"trankilou.fr/lassistanoque/backend/internal/domain"
|
||||
)
|
||||
|
||||
type TursoAgentRepository struct {
|
||||
db *sql.DB
|
||||
AgentTable orm.Table[domain.Agent]
|
||||
}
|
||||
|
||||
func NewTursoAgentRepository(db *sql.DB) *TursoAgentRepository {
|
||||
dialect := orm.NewSqliteDialect()
|
||||
AgentTable, err := orm.NewTable[domain.Agent](db, dialect, "agents")
|
||||
if err != nil {
|
||||
log.Fatalf("error creating lasebuche agent table")
|
||||
}
|
||||
|
||||
return &TursoAgentRepository{
|
||||
db: db,
|
||||
AgentTable: AgentTable,
|
||||
}
|
||||
}
|
||||
|
||||
func (r *TursoAgentRepository) ListAgents(userID string, teamID string) ([]*domain.Agent, error) {
|
||||
return r.AgentTable.Select(
|
||||
orm.WithWhere(
|
||||
"team_id=$1 and team_id in (select team_id from user_teams where user_id=$2)",
|
||||
teamID,
|
||||
userID,
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
func (r *TursoAgentRepository) GetAgent(userID string, teamID string, id string) (*domain.Agent, error) {
|
||||
return r.AgentTable.SelectOne(
|
||||
"id=$1 and team_id=$2 and team_id in (select team_id from user_teams where user_id=$3)",
|
||||
id,
|
||||
teamID,
|
||||
userID,
|
||||
)
|
||||
}
|
||||
|
||||
func (r *TursoAgentRepository) CreateAgent(userID string, agent *domain.Agent) (*domain.Agent, error) {
|
||||
return r.AgentTable.Insert(agent)
|
||||
}
|
||||
|
||||
func (r *TursoAgentRepository) UpdateAgent(userID string, agent *domain.Agent) (*domain.Agent, error) {
|
||||
return r.AgentTable.Update(agent)
|
||||
}
|
||||
|
||||
func (r *TursoAgentRepository) DeleteAgent(userID string, teamID string, id string) error {
|
||||
return r.AgentTable.DeleteWhere(
|
||||
"id=$1 and team_id=$2 and team_id in (select team_id from user_teams where user_id=$3)",
|
||||
id,
|
||||
teamID,
|
||||
userID,
|
||||
)
|
||||
}
|
||||
@@ -69,5 +69,9 @@ func (db *TursoDB) FileRepository() domain.FileRepository {
|
||||
}
|
||||
|
||||
func (db *TursoDB) ProviderRepository() domain.ProviderRepository {
|
||||
return NewTursoModelRepository(db.DB)
|
||||
return NewTursoProviderRepository(db.DB)
|
||||
}
|
||||
|
||||
func (db *TursoDB) AgentRepository() domain.AgentRepository {
|
||||
return NewTursoAgentRepository(db.DB)
|
||||
}
|
||||
@@ -6,19 +6,19 @@ import (
|
||||
"log"
|
||||
"time"
|
||||
|
||||
"gitea.trankilou.fr/fabien/lasebuche"
|
||||
"trankilou.fr/lassistanoque/backend/internal/adapter/database/orm"
|
||||
"trankilou.fr/lassistanoque/backend/internal/domain"
|
||||
"trankilou.fr/lassistanoque/backend/internal/utility"
|
||||
)
|
||||
|
||||
type TursoFileRepository struct {
|
||||
db *sql.DB
|
||||
FileTable lasebuche.Table[domain.File]
|
||||
FileTable orm.Table[domain.File]
|
||||
}
|
||||
|
||||
func NewTursoFileRepository(db *sql.DB) *TursoFileRepository {
|
||||
dialect := lasebuche.NewSqliteDialect()
|
||||
fileTable, err := lasebuche.NewTable[domain.File](db, dialect, "settings")
|
||||
dialect := orm.NewSqliteDialect()
|
||||
fileTable, err := orm.NewTable[domain.File](db, dialect, "settings")
|
||||
if err != nil {
|
||||
log.Fatalf("error creating lasebuche team table")
|
||||
}
|
||||
|
||||
@@ -11,3 +11,4 @@ drop table tools;
|
||||
drop table tasks;
|
||||
drop table history;
|
||||
drop table files;
|
||||
drop table agents;
|
||||
@@ -141,3 +141,24 @@ create table files (
|
||||
_date_updated numeric ,
|
||||
_version text not null
|
||||
);
|
||||
|
||||
create table agents (
|
||||
id text not null primary key,
|
||||
team_id text not null,
|
||||
name text not null default '',
|
||||
system_prompt text not null default '',
|
||||
def_provider_id text not null default '',
|
||||
def_model_id text not null default '',
|
||||
tools_policy text not null default '',
|
||||
sub_agents numeric not null default 0,
|
||||
loop_strategy text not null default '',
|
||||
max_iterations numeric not null default 1,
|
||||
stopping_criteria text not null default '',
|
||||
-- kb text not null default '',
|
||||
-- tools text not null default '',
|
||||
-- skills text not null default '',
|
||||
-- channels text not null default '',
|
||||
_date_created numeric not null default current_timestamp,
|
||||
_date_updated numeric ,
|
||||
_version text not null
|
||||
);
|
||||
@@ -4,37 +4,40 @@ import (
|
||||
"database/sql"
|
||||
"log"
|
||||
|
||||
"gitea.trankilou.fr/fabien/lasebuche"
|
||||
"trankilou.fr/lassistanoque/backend/internal/adapter/database/orm"
|
||||
"trankilou.fr/lassistanoque/backend/internal/domain"
|
||||
)
|
||||
|
||||
type TursoModelRepository struct {
|
||||
type TursoProviderRepository struct {
|
||||
db *sql.DB
|
||||
providerTable lasebuche.Table[domain.Provider]
|
||||
providerTable orm.Table[domain.Provider]
|
||||
}
|
||||
|
||||
func NewTursoModelRepository(db *sql.DB) *TursoModelRepository {
|
||||
dialect := lasebuche.NewSqliteDialect()
|
||||
providerTable, err := lasebuche.NewTable[domain.Provider](db, dialect, "providers")
|
||||
func NewTursoProviderRepository(db *sql.DB) *TursoProviderRepository {
|
||||
dialect := orm.NewSqliteDialect()
|
||||
providerTable, err := orm.NewTable[domain.Provider](db, dialect, "providers")
|
||||
if err != nil {
|
||||
log.Fatalf("error creating lasebuche team table")
|
||||
log.Fatalf("error creating lasebuche provider table")
|
||||
}
|
||||
|
||||
return &TursoModelRepository{
|
||||
return &TursoProviderRepository{
|
||||
db: db,
|
||||
providerTable: providerTable,
|
||||
}
|
||||
}
|
||||
|
||||
func (r *TursoModelRepository) ListProviders(userID string, teamID string) ([]*domain.Provider, error) {
|
||||
return r.providerTable.SelectWhere(
|
||||
"team_id=$1 and team_id in (select team_id from user_teams where user_id=$2)",
|
||||
teamID,
|
||||
userID,
|
||||
func (r *TursoProviderRepository) ListProviders(userID string, teamID string) ([]*domain.Provider, error) {
|
||||
return r.providerTable.Select(
|
||||
orm.WithWhere(
|
||||
"team_id=$1 and team_id in (select team_id from user_teams where user_id=$2)",
|
||||
teamID,
|
||||
userID,
|
||||
),
|
||||
orm.WithOrder("name"),
|
||||
)
|
||||
}
|
||||
|
||||
func (r *TursoModelRepository) GetProvider(userID string, teamID string, id string) (*domain.Provider, error) {
|
||||
func (r *TursoProviderRepository) GetProvider(userID string, teamID string, id string) (*domain.Provider, error) {
|
||||
return r.providerTable.SelectOne(
|
||||
"id=$1 and team_id=$2 and team_id in (select team_id from user_teams where user_id=$3)",
|
||||
id,
|
||||
@@ -43,15 +46,15 @@ func (r *TursoModelRepository) GetProvider(userID string, teamID string, id stri
|
||||
)
|
||||
}
|
||||
|
||||
func (r *TursoModelRepository) CreateProvider(userID string, provider *domain.Provider) (*domain.Provider, error) {
|
||||
func (r *TursoProviderRepository) CreateProvider(userID string, provider *domain.Provider) (*domain.Provider, error) {
|
||||
return r.providerTable.Insert(provider)
|
||||
}
|
||||
|
||||
func (r *TursoModelRepository) UpdateProvider(userID string, provider *domain.Provider) (*domain.Provider, error) {
|
||||
func (r *TursoProviderRepository) UpdateProvider(userID string, provider *domain.Provider) (*domain.Provider, error) {
|
||||
return r.providerTable.Update(provider)
|
||||
}
|
||||
|
||||
func (r *TursoModelRepository) DeleteProvider(userID string, teamID string, id string) error {
|
||||
func (r *TursoProviderRepository) DeleteProvider(userID string, teamID string, id string) error {
|
||||
return r.providerTable.DeleteWhere(
|
||||
"id=$1 and team_id=$2 and team_id in (select team_id from user_teams where user_id=$3)",
|
||||
id,
|
||||
|
||||
@@ -4,18 +4,18 @@ import (
|
||||
"database/sql"
|
||||
"log"
|
||||
|
||||
"gitea.trankilou.fr/fabien/lasebuche"
|
||||
"trankilou.fr/lassistanoque/backend/internal/adapter/database/orm"
|
||||
"trankilou.fr/lassistanoque/backend/internal/domain"
|
||||
)
|
||||
|
||||
type TursoSettingsRepository struct {
|
||||
db *sql.DB
|
||||
SettingsTable lasebuche.Table[domain.Settings]
|
||||
SettingsTable orm.Table[domain.Settings]
|
||||
}
|
||||
|
||||
func NewTursoSettingsRepository(db *sql.DB) *TursoSettingsRepository {
|
||||
dialect := lasebuche.NewSqliteDialect()
|
||||
settingsTable, err := lasebuche.NewTable[domain.Settings](db, dialect, "settings")
|
||||
dialect := orm.NewSqliteDialect()
|
||||
settingsTable, err := orm.NewTable[domain.Settings](db, dialect, "settings")
|
||||
if err != nil {
|
||||
log.Fatalf("error creating lasebuche team table")
|
||||
}
|
||||
|
||||
@@ -4,34 +4,33 @@ import (
|
||||
"database/sql"
|
||||
"log"
|
||||
|
||||
"trankilou.fr/lassistanoque/backend/internal/adapter/database/orm"
|
||||
"trankilou.fr/lassistanoque/backend/internal/domain"
|
||||
|
||||
"gitea.trankilou.fr/fabien/lasebuche"
|
||||
)
|
||||
|
||||
type TursoUserRepository struct {
|
||||
DB *sql.DB
|
||||
UserTable lasebuche.Table[domain.User]
|
||||
TeamTable lasebuche.Table[domain.Team]
|
||||
UserTeamTable lasebuche.Table[domain.UserTeam]
|
||||
AddressTable lasebuche.Table[domain.UserAddress]
|
||||
UserTable orm.Table[domain.User]
|
||||
TeamTable orm.Table[domain.Team]
|
||||
UserTeamTable orm.Table[domain.UserTeam]
|
||||
AddressTable orm.Table[domain.UserAddress]
|
||||
}
|
||||
|
||||
func NewTursoUserRepository(db *sql.DB) *TursoUserRepository {
|
||||
dialect := lasebuche.NewSqliteDialect()
|
||||
userTable, err := lasebuche.NewTable[domain.User](db, dialect, "users")
|
||||
dialect := orm.NewSqliteDialect()
|
||||
userTable, err := orm.NewTable[domain.User](db, dialect, "users")
|
||||
if err != nil {
|
||||
log.Fatalf("error creating lasebuche user table")
|
||||
}
|
||||
teamTable, err := lasebuche.NewTable[domain.Team](db, dialect, "teams")
|
||||
teamTable, err := orm.NewTable[domain.Team](db, dialect, "teams")
|
||||
if err != nil {
|
||||
log.Fatalf("error creating lasebuche team table")
|
||||
}
|
||||
userTeamTable, err := lasebuche.NewTable[domain.UserTeam](db, dialect, "user_teams")
|
||||
userTeamTable, err := orm.NewTable[domain.UserTeam](db, dialect, "user_teams")
|
||||
if err != nil {
|
||||
log.Fatalf("error creating lasebuche team table")
|
||||
}
|
||||
userAddressTable, err := lasebuche.NewTable[domain.UserAddress](db, dialect, "user_addresses")
|
||||
userAddressTable, err := orm.NewTable[domain.UserAddress](db, dialect, "user_addresses")
|
||||
if err != nil {
|
||||
log.Fatalf("error creating lasebuche team table")
|
||||
}
|
||||
@@ -53,7 +52,7 @@ func (ur *TursoUserRepository) FindUserByEmail(email string) (*domain.User, erro
|
||||
}
|
||||
|
||||
func (ur *TursoUserRepository) ListUsers() ([]*domain.User, error) {
|
||||
return ur.UserTable.SelectWhere("")
|
||||
return ur.UserTable.Select()
|
||||
}
|
||||
|
||||
func (ur *TursoUserRepository) CreateUser(user *domain.User) (*domain.User, error) {
|
||||
@@ -73,7 +72,10 @@ func (ur *TursoUserRepository) FindTeam(userid string, teamid string) (*domain.T
|
||||
}
|
||||
|
||||
func (ur *TursoUserRepository) ListTeams(userid string) ([]*domain.Team, error) {
|
||||
return ur.TeamTable.SelectWhere("id in (select team_id from user_teams where user_id=$1)", userid)
|
||||
return ur.TeamTable.Select(
|
||||
orm.WithWhere("id in (select team_id from user_teams where user_id=$1)", userid),
|
||||
orm.WithOrder("label asc"),
|
||||
)
|
||||
}
|
||||
|
||||
func (ur *TursoUserRepository) CreateTeam(userid string, team *domain.Team) (*domain.Team, error) {
|
||||
@@ -108,7 +110,9 @@ func (ur *TursoUserRepository) FindUserTeam(userid string, teamid string) (*doma
|
||||
}
|
||||
|
||||
func (ur *TursoUserRepository) ListUserTeams(userid string) ([]*domain.UserTeam, error) {
|
||||
return ur.UserTeamTable.SelectWhere("user_id=$1", userid)
|
||||
return ur.UserTeamTable.Select(
|
||||
orm.WithWhere("user_id=$1", userid),
|
||||
)
|
||||
}
|
||||
|
||||
func (ur *TursoUserRepository) CreateUserTeam(userid string, team *domain.UserTeam) (*domain.UserTeam, error) {
|
||||
@@ -124,7 +128,10 @@ func (ur *TursoUserRepository) DeleteUserTeam(userid string, teamid string) erro
|
||||
}
|
||||
|
||||
func (ur *TursoUserRepository) ListUserAddresses(id string) ([]*domain.UserAddress, error) {
|
||||
return ur.AddressTable.SelectWhere("user_id=$1", id)
|
||||
return ur.AddressTable.Select(
|
||||
orm.WithWhere("user_id=$1", id),
|
||||
orm.WithOrder("type asc"),
|
||||
)
|
||||
}
|
||||
|
||||
func (ur *TursoUserRepository) GetUserAddress(addrID string) (*domain.UserAddress, error) {
|
||||
|
||||
@@ -12,11 +12,11 @@ import (
|
||||
)
|
||||
|
||||
var providerTypes = []domain.Item{
|
||||
{ID: "anthropic", Text: "Anthropic"},
|
||||
{ID: "openai", Text: "OpenAI"},
|
||||
{ID: "ollama", Text: "Ollama"},
|
||||
{ID: "openaicomp", Text: "OpenAI compatible"},
|
||||
{ID: "openrouter", Text: "Openrouter"},
|
||||
{Value: "anthropic", Label: "Anthropic"},
|
||||
{Value: "openai", Label: "OpenAI"},
|
||||
{Value: "ollama", Label: "Ollama"},
|
||||
{Value: "openaicomp", Label: "OpenAI compatible"},
|
||||
{Value: "openrouter", Label: "Openrouter"},
|
||||
}
|
||||
|
||||
type AnyLLMEngine struct {
|
||||
|
||||
@@ -0,0 +1,28 @@
|
||||
package domain
|
||||
|
||||
import "time"
|
||||
|
||||
type Agent struct {
|
||||
ID string `db:"id" json:"id"`
|
||||
TeamID string `db:"team_id" json:"teamId"`
|
||||
Name string `db:"name" json:"name"`
|
||||
SystemPrompt string `db:"system_prompt" json:"systemPrompt"`
|
||||
DefaultProviderID string `db:"def_provider_id" json:"defaultProviderId"`
|
||||
DefaultModelID string `db:"def_model_id" json:"defaultModelId"`
|
||||
ToolsPolicy string `db:"tools_policy" json:"toolsPolicy"`
|
||||
SubAgents bool `db:"sub_agents" json:"subAgents"`
|
||||
LoopStrategy string `db:"loop_strategy" json:"loopStrategy"`
|
||||
MaxIterations string `db:"max_iterations" json:"maxIterations"`
|
||||
StoppingCriteria string `db:"stopping_criteria" json:"stoppingCriteria"`
|
||||
DateCreated time.Time `db:"_date_created" json:"_date_created"`
|
||||
DateUpdated *time.Time `db:"_date_updated" json:"_date_updated"`
|
||||
VersionId string `db:"_version" json:"_version"`
|
||||
}
|
||||
|
||||
type AgentRepository interface {
|
||||
ListAgents(userID string, teamID string) ([]*Agent, error)
|
||||
GetAgent(userID string, teamID string, id string) (*Agent, error)
|
||||
CreateAgent(userID string, agent *Agent) (*Agent, error)
|
||||
UpdateAgent(userID string, agent *Agent) (*Agent, error)
|
||||
DeleteAgent(userID string, teamID string, id string) error
|
||||
}
|
||||
@@ -1,6 +1,12 @@
|
||||
package domain
|
||||
|
||||
type Item struct {
|
||||
ID string `json:"id"`
|
||||
Text string `json:"text"`
|
||||
Value string `json:"value"`
|
||||
Label string `json:"label"`
|
||||
}
|
||||
|
||||
type ComplexItem struct {
|
||||
Value string `json:"value"`
|
||||
Label string `json:"label"`
|
||||
Group string `json:"group"`
|
||||
}
|
||||
@@ -0,0 +1,100 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/labstack/echo/v5"
|
||||
"trankilou.fr/lassistanoque/backend/internal/domain"
|
||||
"trankilou.fr/lassistanoque/backend/internal/service/agent"
|
||||
"trankilou.fr/lassistanoque/backend/internal/service/auth"
|
||||
)
|
||||
|
||||
func NewAgentGroup(prefix string, e *echo.Group, service *agent.Service, middlewares ...echo.MiddlewareFunc) *echo.Group {
|
||||
agentHandler := &AgentHandler{
|
||||
agentService: service,
|
||||
}
|
||||
agent := e.Group(prefix, middlewares...)
|
||||
|
||||
agent.GET("/:space", agentHandler.ListAgents)
|
||||
agent.GET("/:space/:agent", agentHandler.GetAgent)
|
||||
agent.PUT("/:space", agentHandler.UpdateAgent)
|
||||
agent.POST("/:space", agentHandler.CreateAgent)
|
||||
agent.DELETE("/:space/:agent", agentHandler.DeleteAgent)
|
||||
|
||||
return agent
|
||||
}
|
||||
|
||||
type AgentHandler struct {
|
||||
agentService *agent.Service
|
||||
}
|
||||
|
||||
func (h *AgentHandler) ListAgents(c *echo.Context) error {
|
||||
userID := c.Get(auth.ContextUserIDKey).(string)
|
||||
teamID := c.Param("space")
|
||||
agents, err := h.agentService.ListAgents(userID, teamID)
|
||||
if err != nil {
|
||||
c.Logger().Error("error listing agents", err)
|
||||
return echo.NewHTTPError(http.StatusBadRequest, "error listing agents")
|
||||
}
|
||||
return c.JSON(http.StatusOK, agents)
|
||||
}
|
||||
|
||||
func (h *AgentHandler) GetAgent(c *echo.Context) error {
|
||||
userID := c.Get(auth.ContextUserIDKey).(string)
|
||||
teamID := c.Param("space")
|
||||
agentID := c.Param("agent")
|
||||
agent, err := h.agentService.GetAgent(userID, teamID, agentID)
|
||||
if err != nil {
|
||||
c.Logger().Error("error getting agent: %s", err)
|
||||
return echo.NewHTTPError(http.StatusBadRequest, "error getting agent")
|
||||
}
|
||||
return c.JSON(http.StatusOK, agent)
|
||||
}
|
||||
|
||||
func (h *AgentHandler) CreateAgent(c *echo.Context) error {
|
||||
userID := c.Get(auth.ContextUserIDKey).(string)
|
||||
teamID := c.Param("space")
|
||||
var agent domain.Agent
|
||||
err := c.Bind(&agent)
|
||||
if err != nil {
|
||||
c.Logger().Error("error binding agent : %s", err)
|
||||
return echo.NewHTTPError(http.StatusBadRequest, "error binding agent")
|
||||
}
|
||||
agent.TeamID = teamID
|
||||
updagent, err := h.agentService.CreateAgent(userID, &agent)
|
||||
if err != nil {
|
||||
c.Logger().Error("error creating agent: %s", err)
|
||||
return echo.NewHTTPError(http.StatusBadRequest, "error creating agent")
|
||||
}
|
||||
return c.JSON(http.StatusOK, updagent)
|
||||
}
|
||||
|
||||
func (h *AgentHandler) UpdateAgent(c *echo.Context) error {
|
||||
userID := c.Get(auth.ContextUserIDKey).(string)
|
||||
teamID := c.Param("space")
|
||||
var agent domain.Agent
|
||||
err := c.Bind(&agent)
|
||||
if err != nil {
|
||||
c.Logger().Error("error binding agent : %s", err)
|
||||
return echo.NewHTTPError(http.StatusBadRequest, "error binding agent")
|
||||
}
|
||||
agent.TeamID = teamID
|
||||
updagent, err := h.agentService.UpdateAgent(userID, &agent)
|
||||
if err != nil {
|
||||
c.Logger().Error("error updating agent: %s", err)
|
||||
return echo.NewHTTPError(http.StatusBadRequest, "error updating agent")
|
||||
}
|
||||
return c.JSON(http.StatusOK, updagent)
|
||||
}
|
||||
|
||||
func (h *AgentHandler) DeleteAgent(c *echo.Context) error {
|
||||
userID := c.Get(auth.ContextUserIDKey).(string)
|
||||
teamID := c.Param("space")
|
||||
agentID := c.Param("agent")
|
||||
err := h.agentService.DeleteAgent(userID, teamID, agentID)
|
||||
if err != nil {
|
||||
c.Logger().Error("error creating agent: %s", err)
|
||||
return echo.NewHTTPError(http.StatusBadRequest, "error creating agent")
|
||||
}
|
||||
return c.JSON(http.StatusOK, agentID)
|
||||
}
|
||||
@@ -17,6 +17,7 @@ func NewModelGroup(prefix string, e *echo.Group, service *provider.Service, midd
|
||||
|
||||
model.GET("/providerTypes", modelHandler.ListProviderTypes)
|
||||
model.GET("/:space", modelHandler.ListProviders)
|
||||
model.GET("/:space/models", modelHandler.ListProvidersModels)
|
||||
model.POST("/:space/avalable-models", modelHandler.ListAvailableModels)
|
||||
model.GET("/:space/:provider", modelHandler.GetProvider)
|
||||
model.PUT("/:space", modelHandler.UpdateProvider)
|
||||
@@ -116,8 +117,19 @@ func (h *ModelHandler) ListAvailableModels(c *echo.Context) error {
|
||||
provider.TeamID = teamID
|
||||
list, err := h.providerService.ListAvailableModels(c.Request().Context(), &provider)
|
||||
if err != nil {
|
||||
c.Logger().Error("error listing models: %s", err)
|
||||
return echo.NewHTTPError(http.StatusBadRequest, "error creating provider")
|
||||
c.Logger().Error("error listing available models: %s", err)
|
||||
return echo.NewHTTPError(http.StatusBadRequest, "error listing available models")
|
||||
}
|
||||
return c.JSON(http.StatusOK, list)
|
||||
}
|
||||
|
||||
func (h *ModelHandler) ListProvidersModels(c *echo.Context) error {
|
||||
userID := c.Get(auth.ContextUserIDKey).(string)
|
||||
teamID := c.Param("space")
|
||||
list, err := h.providerService.ListProvidersModels(userID, teamID)
|
||||
if err != nil {
|
||||
c.Logger().Error("error listing providers models: %s", err)
|
||||
return echo.NewHTTPError(http.StatusBadRequest, "error listing providers models")
|
||||
}
|
||||
return c.JSON(http.StatusOK, list)
|
||||
}
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"github.com/labstack/echo/v5/middleware"
|
||||
"trankilou.fr/lassistanoque/backend/internal/config"
|
||||
"trankilou.fr/lassistanoque/backend/internal/http/handlers"
|
||||
"trankilou.fr/lassistanoque/backend/internal/service/agent"
|
||||
"trankilou.fr/lassistanoque/backend/internal/service/auth"
|
||||
"trankilou.fr/lassistanoque/backend/internal/service/provider"
|
||||
"trankilou.fr/lassistanoque/backend/internal/service/storage"
|
||||
@@ -37,6 +38,7 @@ type Dependencies struct {
|
||||
AuthService *auth.Service
|
||||
UserService *user.Service
|
||||
ProviderService *provider.Service
|
||||
AgentService *agent.Service
|
||||
TokenManager auth.TokenManager
|
||||
}
|
||||
|
||||
@@ -74,6 +76,7 @@ func NewRouter(deps Dependencies) *Router {
|
||||
_ = handlers.NewUserGroup("/user", api, deps.UserService, deps.TokenManager.TokenMiddleware)
|
||||
_ = handlers.NewMiscGroup("/misc", api, deps.TokenManager.TokenMiddleware)
|
||||
_ = handlers.NewModelGroup("/provider", api, deps.ProviderService, deps.TokenManager.TokenMiddleware)
|
||||
_ = handlers.NewAgentGroup("/agent", api, deps.AgentService, deps.TokenManager.TokenMiddleware)
|
||||
|
||||
return &Router{
|
||||
echo: e,
|
||||
|
||||
@@ -0,0 +1,64 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"trankilou.fr/lassistanoque/backend/internal/domain"
|
||||
)
|
||||
|
||||
type Service struct {
|
||||
repo domain.AgentRepository
|
||||
repoUser domain.UserRepository
|
||||
}
|
||||
|
||||
func NewService(
|
||||
repo domain.AgentRepository,
|
||||
repoUser domain.UserRepository,
|
||||
) *Service {
|
||||
return &Service{
|
||||
repo,
|
||||
repoUser,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) ListAgents(userID string, teamID string) ([]*domain.Agent, error) {
|
||||
return s.repo.ListAgents(userID, teamID)
|
||||
}
|
||||
|
||||
func (s *Service) GetAgent(userID string, teamID string, id string) (*domain.Agent, error) {
|
||||
return s.repo.GetAgent(userID, teamID, id)
|
||||
}
|
||||
|
||||
func (s *Service) CreateAgent(userID string, Agent *domain.Agent) (*domain.Agent, error) {
|
||||
|
||||
if _, err := s.repoUser.FindUserTeam(userID, Agent.TeamID); err != nil {
|
||||
return nil, fmt.Errorf("error finding user in team: %s", err)
|
||||
}
|
||||
|
||||
return s.repo.CreateAgent(userID, Agent)
|
||||
}
|
||||
|
||||
func (s *Service) UpdateAgent(userID string, Agent *domain.Agent) (*domain.Agent, error) {
|
||||
if _, err := s.repoUser.FindUserTeam(userID, Agent.TeamID); err != nil {
|
||||
return nil, fmt.Errorf("error finding user in team: %s", err)
|
||||
}
|
||||
|
||||
if _, err := s.repo.GetAgent(userID, Agent.TeamID, Agent.ID); err != nil {
|
||||
return nil, fmt.Errorf("error finding Agent: %s", err)
|
||||
}
|
||||
|
||||
return s.repo.UpdateAgent(userID, Agent)
|
||||
}
|
||||
|
||||
func (s *Service) DeleteAgent(userID string, teamID string, id string) error {
|
||||
if _, err := s.repoUser.FindUserTeam(userID, teamID); err != nil {
|
||||
return fmt.Errorf("error finding user in team: %s", err)
|
||||
}
|
||||
|
||||
if _, err := s.repo.GetAgent(userID, teamID, id); err != nil {
|
||||
return fmt.Errorf("error finding Agent: %s", err)
|
||||
}
|
||||
|
||||
return s.repo.DeleteAgent(userID, teamID, id)
|
||||
|
||||
}
|
||||
@@ -3,6 +3,7 @@ package provider
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"trankilou.fr/lassistanoque/backend/internal/domain"
|
||||
)
|
||||
@@ -33,6 +34,28 @@ func (s *Service) ListProviders(userID string, teamID string) ([]*domain.Provide
|
||||
return s.repo.ListProviders(userID, teamID)
|
||||
}
|
||||
|
||||
func (s *Service) ListProvidersModels(userID string, teamID string) ([]*domain.ComplexItem, error) {
|
||||
|
||||
items := make([]*domain.ComplexItem, 0)
|
||||
|
||||
providers, err := s.repo.ListProviders(userID, teamID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, provider := range providers {
|
||||
models := strings.Split(provider.Models, "|")
|
||||
for _, m := range models {
|
||||
items = append(items, &domain.ComplexItem{
|
||||
Group: provider.Name,
|
||||
Value: provider.ID + "|" + m,
|
||||
Label: m,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func (s *Service) GetProvider(userID string, teamID string, id string) (*domain.Provider, error) {
|
||||
return s.repo.GetProvider(userID, teamID, id)
|
||||
}
|
||||
|
||||
Reference in new issue
Block a user