Initial commit of Descrybe v2 without local scratch artifacts.
Drop one-shot tmp/axe scripts and agent i18n scratch so the Gitea tree is deployable.
This commit is contained in:
@@ -0,0 +1,211 @@
|
||||
package company
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/descrybe/descrybe-v2/apps/api/internal/security"
|
||||
"github.com/google/uuid"
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
)
|
||||
|
||||
// BrandKit stores company brand voice, guidelines, and visual identity.
|
||||
type BrandKit struct {
|
||||
CompanyID uuid.UUID `json:"company_id"`
|
||||
VoiceTone string `json:"voice_tone"`
|
||||
Dos []string `json:"dos"`
|
||||
Donts []string `json:"donts"`
|
||||
PrimaryColor string `json:"primary_color"`
|
||||
SecondaryColor string `json:"secondary_color"`
|
||||
LogoURL string `json:"logo_url"`
|
||||
PreferredTerms []string `json:"preferred_terms"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
// EmptyBrand returns a zero kit for a company (no row yet).
|
||||
func EmptyBrand(companyID uuid.UUID) BrandKit {
|
||||
return BrandKit{
|
||||
CompanyID: companyID,
|
||||
Dos: []string{},
|
||||
Donts: []string{},
|
||||
PreferredTerms: []string{},
|
||||
}
|
||||
}
|
||||
|
||||
// LoadBrand returns the company brand kit, or an empty kit when none is saved.
|
||||
func LoadBrand(ctx context.Context, pool *pgxpool.Pool, companyID uuid.UUID) (BrandKit, error) {
|
||||
var b BrandKit
|
||||
err := pool.QueryRow(ctx, `
|
||||
SELECT company_id, voice_tone, COALESCE(dos, '{}'), COALESCE(donts, '{}'),
|
||||
primary_color, secondary_color, logo_url, COALESCE(preferred_terms, '{}'), updated_at
|
||||
FROM company_brand WHERE company_id = $1`, companyID).
|
||||
Scan(&b.CompanyID, &b.VoiceTone, &b.Dos, &b.Donts,
|
||||
&b.PrimaryColor, &b.SecondaryColor, &b.LogoURL, &b.PreferredTerms, &b.UpdatedAt)
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return EmptyBrand(companyID), nil
|
||||
}
|
||||
if err != nil {
|
||||
return BrandKit{}, err
|
||||
}
|
||||
b.Dos = cleanStrings(b.Dos)
|
||||
b.Donts = cleanStrings(b.Donts)
|
||||
b.PreferredTerms = cleanStrings(b.PreferredTerms)
|
||||
return b, nil
|
||||
}
|
||||
|
||||
// UpsertBrand saves the brand kit for a company.
|
||||
func UpsertBrand(ctx context.Context, pool *pgxpool.Pool, companyID uuid.UUID, in BrandKit) (BrandKit, error) {
|
||||
in.VoiceTone = security.SanitizePrompt(in.VoiceTone, security.MaxBrandFieldRunes)
|
||||
in.PrimaryColor = security.TruncateRunes(strings.TrimSpace(in.PrimaryColor), 32)
|
||||
in.SecondaryColor = security.TruncateRunes(strings.TrimSpace(in.SecondaryColor), 32)
|
||||
logo, err := ValidateLogoURL(in.LogoURL, companyID)
|
||||
if err != nil {
|
||||
return BrandKit{}, err
|
||||
}
|
||||
in.LogoURL = logo
|
||||
in.Dos = security.SanitizeBrandList(in.Dos)
|
||||
in.Donts = security.SanitizeBrandList(in.Donts)
|
||||
in.PreferredTerms = security.SanitizeBrandList(in.PreferredTerms)
|
||||
|
||||
var b BrandKit
|
||||
err = pool.QueryRow(ctx, `
|
||||
INSERT INTO company_brand (
|
||||
company_id, voice_tone, dos, donts, primary_color, secondary_color, logo_url, preferred_terms, updated_at
|
||||
) VALUES ($1,$2,$3,$4,$5,$6,$7,$8, now())
|
||||
ON CONFLICT (company_id) DO UPDATE SET
|
||||
voice_tone = EXCLUDED.voice_tone,
|
||||
dos = EXCLUDED.dos,
|
||||
donts = EXCLUDED.donts,
|
||||
primary_color = EXCLUDED.primary_color,
|
||||
secondary_color = EXCLUDED.secondary_color,
|
||||
logo_url = EXCLUDED.logo_url,
|
||||
preferred_terms = EXCLUDED.preferred_terms,
|
||||
updated_at = now()
|
||||
RETURNING company_id, voice_tone, COALESCE(dos, '{}'), COALESCE(donts, '{}'),
|
||||
primary_color, secondary_color, logo_url, COALESCE(preferred_terms, '{}'), updated_at`,
|
||||
companyID, in.VoiceTone, in.Dos, in.Donts, in.PrimaryColor, in.SecondaryColor, in.LogoURL, in.PreferredTerms,
|
||||
).Scan(&b.CompanyID, &b.VoiceTone, &b.Dos, &b.Donts,
|
||||
&b.PrimaryColor, &b.SecondaryColor, &b.LogoURL, &b.PreferredTerms, &b.UpdatedAt)
|
||||
if err != nil {
|
||||
return BrandKit{}, err
|
||||
}
|
||||
if b.Dos == nil {
|
||||
b.Dos = []string{}
|
||||
}
|
||||
if b.Donts == nil {
|
||||
b.Donts = []string{}
|
||||
}
|
||||
if b.PreferredTerms == nil {
|
||||
b.PreferredTerms = []string{}
|
||||
}
|
||||
return b, nil
|
||||
}
|
||||
|
||||
// HasContent reports whether any brand guidance is configured.
|
||||
func (b BrandKit) HasContent() bool {
|
||||
return strings.TrimSpace(b.VoiceTone) != "" ||
|
||||
len(b.Dos) > 0 ||
|
||||
len(b.Donts) > 0 ||
|
||||
len(b.PreferredTerms) > 0 ||
|
||||
strings.TrimSpace(b.PrimaryColor) != "" ||
|
||||
strings.TrimSpace(b.SecondaryColor) != "" ||
|
||||
strings.TrimSpace(b.LogoURL) != ""
|
||||
}
|
||||
|
||||
// PromptBlock formats brand voice instructions for AI system prompts.
|
||||
// Returns empty string when the kit has no usable voice content.
|
||||
// Kept short (bullet lines) for weak local models / 8k context.
|
||||
func (b BrandKit) PromptBlock() string {
|
||||
var parts []string
|
||||
if t := security.SanitizePrompt(b.VoiceTone, 160); t != "" {
|
||||
parts = append(parts, "- tone: "+t)
|
||||
}
|
||||
dos := security.SanitizeBrandList(b.Dos)
|
||||
donts := security.SanitizeBrandList(b.Donts)
|
||||
terms := security.SanitizeBrandList(b.PreferredTerms)
|
||||
if len(dos) > 4 {
|
||||
dos = dos[:4]
|
||||
}
|
||||
if len(donts) > 4 {
|
||||
donts = donts[:4]
|
||||
}
|
||||
if len(terms) > 6 {
|
||||
terms = terms[:6]
|
||||
}
|
||||
if len(dos) > 0 {
|
||||
parts = append(parts, "- do: "+strings.Join(dos, "; "))
|
||||
}
|
||||
if len(donts) > 0 {
|
||||
parts = append(parts, "- don't: "+strings.Join(donts, "; "))
|
||||
}
|
||||
if len(terms) > 0 {
|
||||
parts = append(parts, "- terms: "+strings.Join(terms, ", "))
|
||||
}
|
||||
if len(parts) == 0 {
|
||||
return ""
|
||||
}
|
||||
block := "Brand:\n" + strings.Join(parts, "\n")
|
||||
return security.TruncateRunes(block, 500)
|
||||
}
|
||||
|
||||
// FormulaTips returns short brand-aware tips for formula/preview UI.
|
||||
func (b BrandKit) FormulaTips() []string {
|
||||
tips := make([]string, 0, 4)
|
||||
if t := strings.TrimSpace(b.VoiceTone); t != "" {
|
||||
tips = append(tips, "Match brand tone: "+truncateTip(t, 120))
|
||||
}
|
||||
if len(b.PreferredTerms) > 0 {
|
||||
n := len(b.PreferredTerms)
|
||||
if n > 5 {
|
||||
n = 5
|
||||
}
|
||||
tips = append(tips, "Prefer terms: "+strings.Join(b.PreferredTerms[:n], ", "))
|
||||
}
|
||||
if len(b.Donts) > 0 {
|
||||
n := len(b.Donts)
|
||||
if n > 3 {
|
||||
n = 3
|
||||
}
|
||||
tips = append(tips, "Avoid: "+strings.Join(b.Donts[:n], "; "))
|
||||
}
|
||||
if len(b.Dos) > 0 {
|
||||
n := len(b.Dos)
|
||||
if n > 3 {
|
||||
n = 3
|
||||
}
|
||||
tips = append(tips, "Do: "+strings.Join(b.Dos[:n], "; "))
|
||||
}
|
||||
return tips
|
||||
}
|
||||
|
||||
func cleanStrings(in []string) []string {
|
||||
if len(in) == 0 {
|
||||
return []string{}
|
||||
}
|
||||
out := make([]string, 0, len(in))
|
||||
seen := map[string]struct{}{}
|
||||
for _, s := range in {
|
||||
s = strings.TrimSpace(s)
|
||||
if s == "" {
|
||||
continue
|
||||
}
|
||||
key := strings.ToLower(s)
|
||||
if _, ok := seen[key]; ok {
|
||||
continue
|
||||
}
|
||||
seen[key] = struct{}{}
|
||||
out = append(out, s)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func truncateTip(s string, max int) string {
|
||||
s = strings.TrimSpace(s)
|
||||
if max <= 0 || len(s) <= max {
|
||||
return s
|
||||
}
|
||||
return strings.TrimSpace(s[:max]) + "…"
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
package company
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestBrandKit_PromptBlock(t *testing.T) {
|
||||
empty := BrandKit{}
|
||||
if empty.PromptBlock() != "" {
|
||||
t.Fatalf("empty should yield empty prompt")
|
||||
}
|
||||
b := BrandKit{
|
||||
VoiceTone: "confident, concise",
|
||||
Dos: []string{"Lead with benefit"},
|
||||
Donts: []string{"No hype"},
|
||||
PreferredTerms: []string{"wireless", "premium"},
|
||||
}
|
||||
got := b.PromptBlock()
|
||||
for _, want := range []string{"Brand:", "confident", "Lead with benefit", "No hype", "wireless"} {
|
||||
if !contains(got, want) {
|
||||
t.Fatalf("prompt missing %q: %s", want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestBrandKit_FormulaTips(t *testing.T) {
|
||||
b := BrandKit{VoiceTone: "warm", PreferredTerms: []string{"eco"}}
|
||||
tips := b.FormulaTips()
|
||||
if len(tips) < 2 {
|
||||
t.Fatalf("tips=%v", tips)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanStringsDedup(t *testing.T) {
|
||||
got := cleanStrings([]string{" A ", "a", "", "B"})
|
||||
if len(got) != 2 || got[0] != "A" || got[1] != "B" {
|
||||
t.Fatalf("got=%v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func contains(s, sub string) bool {
|
||||
return len(s) >= len(sub) && (s == sub || len(sub) == 0 ||
|
||||
(func() bool {
|
||||
for i := 0; i+len(sub) <= len(s); i++ {
|
||||
if s[i:i+len(sub)] == sub {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
})())
|
||||
}
|
||||
@@ -0,0 +1,254 @@
|
||||
package company
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/descrybe/descrybe-v2/apps/api/internal/security"
|
||||
"github.com/google/uuid"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
)
|
||||
|
||||
// LangPromptMap is language-code → prompt text for category / template overrides.
|
||||
type LangPromptMap map[string]string
|
||||
|
||||
// LocalizedFields holds AI/output fields for one content language.
|
||||
type LocalizedFields struct {
|
||||
ProcessedName string `json:"processed_name,omitempty"`
|
||||
ProcessedDescription string `json:"processed_description,omitempty"`
|
||||
MetaTitle string `json:"meta_title,omitempty"`
|
||||
MetaDescription string `json:"meta_description,omitempty"`
|
||||
EnhanceInputHash string `json:"enhance_input_hash,omitempty"`
|
||||
}
|
||||
|
||||
// LocalizedContent is language-code → per-language product output fields.
|
||||
type LocalizedContent map[string]LocalizedFields
|
||||
|
||||
// SanitizeLangPromptMap validates language codes, sanitizes prompts, and drops empties.
|
||||
func SanitizeLangPromptMap(in map[string]string, maxRunes int) (LangPromptMap, error) {
|
||||
out := make(LangPromptMap)
|
||||
if len(in) == 0 {
|
||||
return out, nil
|
||||
}
|
||||
for lang, prompt := range in {
|
||||
code, err := ParseLanguage(lang, false)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("unsupported language %q", lang)
|
||||
}
|
||||
p := strings.TrimSpace(security.SanitizePrompt(prompt, maxRunes))
|
||||
if p == "" {
|
||||
continue
|
||||
}
|
||||
out[code] = p
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// PromptForLanguage returns the prompt for lang, or empty if unset.
|
||||
func PromptForLanguage(m LangPromptMap, lang string) string {
|
||||
if len(m) == 0 {
|
||||
return ""
|
||||
}
|
||||
code, err := ParseLanguage(lang, true)
|
||||
if err != nil {
|
||||
code = DefaultLanguage
|
||||
}
|
||||
return strings.TrimSpace(m[code])
|
||||
}
|
||||
|
||||
// HasAnyPrompt reports whether any language has a non-empty prompt.
|
||||
func HasAnyPrompt(m LangPromptMap) bool {
|
||||
for _, p := range m {
|
||||
if strings.TrimSpace(p) != "" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// DecodeLangPromptMap accepts JSON object / map[string]any / map[string]string.
|
||||
func DecodeLangPromptMap(raw any) (LangPromptMap, error) {
|
||||
out := make(LangPromptMap)
|
||||
if raw == nil {
|
||||
return out, nil
|
||||
}
|
||||
switch v := raw.(type) {
|
||||
case LangPromptMap:
|
||||
return SanitizeLangPromptMap(v, security.MaxCampaignPromptRunes)
|
||||
case map[string]string:
|
||||
return SanitizeLangPromptMap(v, security.MaxCampaignPromptRunes)
|
||||
case map[string]any:
|
||||
tmp := make(map[string]string, len(v))
|
||||
for k, val := range v {
|
||||
s, ok := val.(string)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("prompt for %q must be a string", k)
|
||||
}
|
||||
tmp[k] = s
|
||||
}
|
||||
return SanitizeLangPromptMap(tmp, security.MaxCampaignPromptRunes)
|
||||
case string:
|
||||
s := strings.TrimSpace(v)
|
||||
if s == "" || s == "{}" {
|
||||
return out, nil
|
||||
}
|
||||
var obj map[string]string
|
||||
if err := json.Unmarshal([]byte(s), &obj); err != nil {
|
||||
return nil, fmt.Errorf("invalid prompt map json")
|
||||
}
|
||||
return SanitizeLangPromptMap(obj, security.MaxCampaignPromptRunes)
|
||||
case []byte:
|
||||
if len(v) == 0 {
|
||||
return out, nil
|
||||
}
|
||||
var obj map[string]string
|
||||
if err := json.Unmarshal(v, &obj); err != nil {
|
||||
return nil, fmt.Errorf("invalid prompt map json")
|
||||
}
|
||||
return SanitizeLangPromptMap(obj, security.MaxCampaignPromptRunes)
|
||||
default:
|
||||
b, err := json.Marshal(raw)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid prompt map")
|
||||
}
|
||||
var obj map[string]string
|
||||
if err := json.Unmarshal(b, &obj); err != nil {
|
||||
return nil, fmt.Errorf("invalid prompt map json")
|
||||
}
|
||||
return SanitizeLangPromptMap(obj, security.MaxCampaignPromptRunes)
|
||||
}
|
||||
}
|
||||
|
||||
// EncodeLangPromptMap marshals a prompt map to JSON bytes (never null).
|
||||
func EncodeLangPromptMap(m LangPromptMap) ([]byte, error) {
|
||||
if m == nil {
|
||||
return []byte("{}"), nil
|
||||
}
|
||||
b, err := json.Marshal(m)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return b, nil
|
||||
}
|
||||
|
||||
// ParseContentLanguages validates and normalizes an ordered language list.
|
||||
// Empty input with allowEmptyAsPrimary yields [DefaultLanguage] or [primary] when primary set.
|
||||
func ParseContentLanguages(raw []string, primary string) ([]string, error) {
|
||||
primaryCode, err := ParseLanguage(primary, true)
|
||||
if err != nil {
|
||||
primaryCode = DefaultLanguage
|
||||
}
|
||||
seen := map[string]struct{}{}
|
||||
out := make([]string, 0, len(raw)+1)
|
||||
add := func(code string) {
|
||||
if _, ok := seen[code]; ok {
|
||||
return
|
||||
}
|
||||
seen[code] = struct{}{}
|
||||
out = append(out, code)
|
||||
}
|
||||
add(primaryCode)
|
||||
for _, r := range raw {
|
||||
code, err := ParseLanguage(r, false)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("unsupported language %q", r)
|
||||
}
|
||||
add(code)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// LoadContentLanguages returns companies.content_languages, ensuring primary is first.
|
||||
func LoadContentLanguages(ctx context.Context, pool *pgxpool.Pool, companyID uuid.UUID) []string {
|
||||
primary := LoadLanguage(ctx, pool, companyID)
|
||||
if pool == nil {
|
||||
return []string{primary}
|
||||
}
|
||||
var langs []string
|
||||
err := pool.QueryRow(ctx, `
|
||||
SELECT COALESCE(content_languages, '{}') FROM companies WHERE id = $1`, companyID).Scan(&langs)
|
||||
if err != nil || len(langs) == 0 {
|
||||
return []string{primary}
|
||||
}
|
||||
parsed, err := ParseContentLanguages(langs, primary)
|
||||
if err != nil {
|
||||
return []string{primary}
|
||||
}
|
||||
return parsed
|
||||
}
|
||||
|
||||
// FieldsForLanguage returns localized fields for lang (empty struct if missing).
|
||||
func FieldsForLanguage(content LocalizedContent, lang string) LocalizedFields {
|
||||
if len(content) == 0 {
|
||||
return LocalizedFields{}
|
||||
}
|
||||
code, err := ParseLanguage(lang, true)
|
||||
if err != nil {
|
||||
code = DefaultLanguage
|
||||
}
|
||||
return content[code]
|
||||
}
|
||||
|
||||
// SetFieldsForLanguage upserts fields for one language into content.
|
||||
func SetFieldsForLanguage(content LocalizedContent, lang string, fields LocalizedFields) LocalizedContent {
|
||||
if content == nil {
|
||||
content = LocalizedContent{}
|
||||
}
|
||||
code, err := ParseLanguage(lang, true)
|
||||
if err != nil {
|
||||
code = DefaultLanguage
|
||||
}
|
||||
content[code] = fields
|
||||
return content
|
||||
}
|
||||
|
||||
// DecodeLocalizedContent parses JSONB / map into LocalizedContent.
|
||||
func DecodeLocalizedContent(raw any) (LocalizedContent, error) {
|
||||
out := LocalizedContent{}
|
||||
if raw == nil {
|
||||
return out, nil
|
||||
}
|
||||
var b []byte
|
||||
switch v := raw.(type) {
|
||||
case []byte:
|
||||
b = v
|
||||
case string:
|
||||
b = []byte(v)
|
||||
default:
|
||||
var err error
|
||||
b, err = json.Marshal(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
if len(b) == 0 || string(b) == "null" || string(b) == "{}" {
|
||||
return out, nil
|
||||
}
|
||||
var tmp map[string]LocalizedFields
|
||||
if err := json.Unmarshal(b, &tmp); err != nil {
|
||||
return nil, fmt.Errorf("invalid localized_content")
|
||||
}
|
||||
for lang, fields := range tmp {
|
||||
code, err := ParseLanguage(lang, false)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
out[code] = fields
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// EncodeLocalizedContent marshals localized content (never null).
|
||||
func EncodeLocalizedContent(c LocalizedContent) ([]byte, error) {
|
||||
if c == nil {
|
||||
return []byte("{}"), nil
|
||||
}
|
||||
return json.Marshal(c)
|
||||
}
|
||||
|
||||
// SyncPrimaryFromLocalized copies primary-language fields onto the denormalized columns shape.
|
||||
func SyncPrimaryFromLocalized(content LocalizedContent, primary string) LocalizedFields {
|
||||
return FieldsForLanguage(content, primary)
|
||||
}
|
||||
@@ -0,0 +1,81 @@
|
||||
package company
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestSanitizeLangPromptMap(t *testing.T) {
|
||||
t.Parallel()
|
||||
_, err := SanitizeLangPromptMap(map[string]string{
|
||||
"SL": " hello {{name}} ",
|
||||
"xx": "bad",
|
||||
}, 100)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for unsupported language")
|
||||
}
|
||||
m, err := SanitizeLangPromptMap(map[string]string{
|
||||
"SL": " hello {{name}} ",
|
||||
"en": "",
|
||||
}, 100)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if m["sl"] != "hello {{name}}" {
|
||||
t.Fatalf("got %#v", m)
|
||||
}
|
||||
if _, ok := m["en"]; ok {
|
||||
t.Fatalf("empty en should be dropped: %#v", m)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPromptForLanguage(t *testing.T) {
|
||||
t.Parallel()
|
||||
m := LangPromptMap{"sl": "slo", "en": "eng"}
|
||||
if got := PromptForLanguage(m, "SL"); got != "slo" {
|
||||
t.Fatalf("got %q", got)
|
||||
}
|
||||
if got := PromptForLanguage(m, "de"); got != "" {
|
||||
t.Fatalf("expected empty, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseContentLanguagesPrimaryFirst(t *testing.T) {
|
||||
t.Parallel()
|
||||
got, err := ParseContentLanguages([]string{"en", "de", "sl"}, "sl")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
want := []string{"sl", "en", "de"}
|
||||
if len(got) != len(want) {
|
||||
t.Fatalf("got %#v", got)
|
||||
}
|
||||
for i := range want {
|
||||
if got[i] != want[i] {
|
||||
t.Fatalf("got %#v want %#v", got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestLocalizedContentRoundTrip(t *testing.T) {
|
||||
t.Parallel()
|
||||
c := LocalizedContent{
|
||||
"sl": {ProcessedName: "Naslov", ProcessedDescription: "Opis"},
|
||||
}
|
||||
b, err := EncodeLocalizedContent(c)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var raw any
|
||||
if err := json.Unmarshal(b, &raw); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
decoded, err := DecodeLocalizedContent(raw)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
f := FieldsForLanguage(decoded, "sl")
|
||||
if f.ProcessedName != "Naslov" || f.ProcessedDescription != "Opis" {
|
||||
t.Fatalf("got %#v", f)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,119 @@
|
||||
package company
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
)
|
||||
|
||||
// DefaultLanguage is the content-language fallback when unset.
|
||||
const DefaultLanguage = "en"
|
||||
|
||||
// ContentLanguages is the allowlist for companies.language (AI/product content).
|
||||
// Keep in sync with apps/web/src/lib/content-languages.ts.
|
||||
var ContentLanguages = []string{
|
||||
"en", "fr", "de", "es", "it", "nl", "pt", "pl",
|
||||
"cs", "sk", "hu", "ro", "bg", "hr", "sl",
|
||||
"sv", "da", "fi", "el", "et", "lv", "lt", "mt", "ga",
|
||||
"ja", // CJK plug-in slot (Japanese)
|
||||
}
|
||||
|
||||
// contentLanguageLabels are English display names for AI prompt injection.
|
||||
// Keep in sync with apps/web/src/lib/content-languages.ts labels.
|
||||
var contentLanguageLabels = map[string]string{
|
||||
"en": "English",
|
||||
"fr": "French",
|
||||
"de": "German",
|
||||
"es": "Spanish",
|
||||
"it": "Italian",
|
||||
"nl": "Dutch",
|
||||
"pt": "Portuguese",
|
||||
"pl": "Polish",
|
||||
"cs": "Czech",
|
||||
"sk": "Slovak",
|
||||
"hu": "Hungarian",
|
||||
"ro": "Romanian",
|
||||
"bg": "Bulgarian",
|
||||
"hr": "Croatian",
|
||||
"sl": "Slovenian",
|
||||
"sv": "Swedish",
|
||||
"da": "Danish",
|
||||
"fi": "Finnish",
|
||||
"el": "Greek",
|
||||
"et": "Estonian",
|
||||
"lv": "Latvian",
|
||||
"lt": "Lithuanian",
|
||||
"mt": "Maltese",
|
||||
"ga": "Irish",
|
||||
"ja": "Japanese",
|
||||
}
|
||||
|
||||
var contentLanguageSet map[string]struct{}
|
||||
|
||||
func init() {
|
||||
contentLanguageSet = make(map[string]struct{}, len(ContentLanguages))
|
||||
for _, code := range ContentLanguages {
|
||||
contentLanguageSet[code] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
// NormalizeLanguage trims and lowercases a content-language code.
|
||||
func NormalizeLanguage(raw string) string {
|
||||
return strings.ToLower(strings.TrimSpace(raw))
|
||||
}
|
||||
|
||||
// IsAllowedLanguage reports whether code is in ContentLanguages (after normalize).
|
||||
func IsAllowedLanguage(raw string) bool {
|
||||
_, ok := contentLanguageSet[NormalizeLanguage(raw)]
|
||||
return ok
|
||||
}
|
||||
|
||||
// ParseLanguage validates and normalizes a content-language code.
|
||||
// Empty input returns DefaultLanguage when allowEmptyAsDefault is true;
|
||||
// otherwise empty is an error (use for explicit PATCH language fields).
|
||||
func ParseLanguage(raw string, allowEmptyAsDefault bool) (string, error) {
|
||||
code := NormalizeLanguage(raw)
|
||||
if code == "" {
|
||||
if allowEmptyAsDefault {
|
||||
return DefaultLanguage, nil
|
||||
}
|
||||
return "", fmt.Errorf("language is required")
|
||||
}
|
||||
if !IsAllowedLanguage(code) {
|
||||
return "", fmt.Errorf("unsupported language %q", code)
|
||||
}
|
||||
return code, nil
|
||||
}
|
||||
|
||||
// LanguageLabel returns the English display name for a content-language code
|
||||
// (for AI prompt injection). Empty/unknown codes fall back to English.
|
||||
func LanguageLabel(raw string) string {
|
||||
code, err := ParseLanguage(raw, true)
|
||||
if err != nil {
|
||||
code = DefaultLanguage
|
||||
}
|
||||
if label, ok := contentLanguageLabels[code]; ok {
|
||||
return label
|
||||
}
|
||||
return contentLanguageLabels[DefaultLanguage]
|
||||
}
|
||||
|
||||
// LoadLanguage returns companies.language for companyID, or DefaultLanguage.
|
||||
func LoadLanguage(ctx context.Context, pool *pgxpool.Pool, companyID uuid.UUID) string {
|
||||
if pool == nil {
|
||||
return DefaultLanguage
|
||||
}
|
||||
var raw string
|
||||
err := pool.QueryRow(ctx, `SELECT COALESCE(language, '') FROM companies WHERE id = $1`, companyID).Scan(&raw)
|
||||
if err != nil {
|
||||
return DefaultLanguage
|
||||
}
|
||||
code, err := ParseLanguage(raw, true)
|
||||
if err != nil {
|
||||
return DefaultLanguage
|
||||
}
|
||||
return code
|
||||
}
|
||||
@@ -0,0 +1,115 @@
|
||||
package company
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestNormalizeLanguage(t *testing.T) {
|
||||
if got := NormalizeLanguage(" EN "); got != "en" {
|
||||
t.Fatalf("NormalizeLanguage: got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseLanguage_AllowedPopular(t *testing.T) {
|
||||
for _, code := range []string{"en", "es", "fr", "de", "it", "pt", "nl", "pl", "ja"} {
|
||||
got, err := ParseLanguage(code, false)
|
||||
if err != nil || got != code {
|
||||
t.Fatalf("ParseLanguage(%q): got=%q err=%v", code, got, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseLanguage_RejectsUnknown(t *testing.T) {
|
||||
if _, err := ParseLanguage("xx", false); err == nil {
|
||||
t.Fatal("expected error for unknown language")
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseLanguage_EmptyDefault(t *testing.T) {
|
||||
got, err := ParseLanguage("", true)
|
||||
if err != nil || got != DefaultLanguage {
|
||||
t.Fatalf("empty default: got=%q err=%v", got, err)
|
||||
}
|
||||
if _, err := ParseLanguage("", false); err == nil {
|
||||
t.Fatal("expected error for empty without default")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsAllowedLanguage(t *testing.T) {
|
||||
if !IsAllowedLanguage("JA") {
|
||||
t.Fatal("ja should be allowed")
|
||||
}
|
||||
if IsAllowedLanguage("zh") {
|
||||
t.Fatal("zh not in allowlist yet")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLanguageLabel(t *testing.T) {
|
||||
if got := LanguageLabel("fr"); got != "French" {
|
||||
t.Fatalf("fr: got %q", got)
|
||||
}
|
||||
if got := LanguageLabel(""); got != "English" {
|
||||
t.Fatalf("empty: got %q", got)
|
||||
}
|
||||
if got := LanguageLabel("xx"); got != "English" {
|
||||
t.Fatalf("unknown: got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestContentLanguagesSyncWithWebAllowlist(t *testing.T) {
|
||||
root := findRepoRoot(t)
|
||||
tsPath := filepath.Join(root, "apps", "web", "src", "lib", "content-languages.ts")
|
||||
raw, err := os.ReadFile(tsPath)
|
||||
if err != nil {
|
||||
t.Fatalf("read %s: %v", tsPath, err)
|
||||
}
|
||||
webCodes := parseTSContentLanguageValues(string(raw))
|
||||
if len(webCodes) == 0 {
|
||||
t.Fatal("no value: \"xx\" entries parsed from content-languages.ts")
|
||||
}
|
||||
if len(webCodes) != len(ContentLanguages) {
|
||||
t.Fatalf("length mismatch: web=%d go=%d\nweb=%v\ngo=%v", len(webCodes), len(ContentLanguages), webCodes, ContentLanguages)
|
||||
}
|
||||
for i, code := range ContentLanguages {
|
||||
if webCodes[i] != code {
|
||||
t.Fatalf("index %d: web=%q go=%q (keep content-languages.ts in sync with ContentLanguages)", i, webCodes[i], code)
|
||||
}
|
||||
if _, err := ParseLanguage(code, false); err != nil {
|
||||
t.Fatalf("ParseLanguage(%q): %v", code, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func findRepoRoot(t *testing.T) string {
|
||||
t.Helper()
|
||||
dir, err := os.Getwd()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for i := 0; i < 10; i++ {
|
||||
candidate := filepath.Join(dir, "apps", "web", "src", "lib", "content-languages.ts")
|
||||
if _, err := os.Stat(candidate); err == nil {
|
||||
return dir
|
||||
}
|
||||
parent := filepath.Dir(dir)
|
||||
if parent == dir {
|
||||
break
|
||||
}
|
||||
dir = parent
|
||||
}
|
||||
t.Fatal("monorepo root not found (expected apps/web/src/lib/content-languages.ts)")
|
||||
return ""
|
||||
}
|
||||
|
||||
func parseTSContentLanguageValues(src string) []string {
|
||||
re := regexp.MustCompile(`value:\s*"([a-z]{2})"`)
|
||||
matches := re.FindAllStringSubmatch(src, -1)
|
||||
out := make([]string, 0, len(matches))
|
||||
for _, m := range matches {
|
||||
out = append(out, m[1])
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,331 @@
|
||||
package company
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/descrybe/descrybe-v2/apps/api/internal/security"
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
const (
|
||||
maxBrandLogoBytes = 2 << 20 // 2 MiB
|
||||
brandLogoSubdir = "brand"
|
||||
// BrandLogoURLPrefix is the authenticated same-origin path stored in logo_url.
|
||||
BrandLogoURLPrefix = "/api/brand/logo/files/"
|
||||
// PublicBrandLogoPathPrefix is the signed public serve path.
|
||||
PublicBrandLogoPathPrefix = "/api/public/brand-logo/"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrLogoInvalidType = errors.New("logo must be PNG, JPEG, or WebP")
|
||||
ErrLogoTooLarge = errors.New("logo exceeds 2 MiB limit")
|
||||
ErrLogoInvalidName = errors.New("invalid logo filename")
|
||||
ErrLogoNotFound = errors.New("logo not found")
|
||||
ErrLogoForbidden = errors.New("logo access forbidden")
|
||||
ErrLogoBadSig = errors.New("invalid or expired logo signature")
|
||||
|
||||
brandLogoNameRE = regexp.MustCompile(`(?i)^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}\.(png|jpe?g|webp)$`)
|
||||
)
|
||||
|
||||
// ClientError reports whether err is a known client-facing brand logo validation error.
|
||||
func ClientError(err error) (msg string, ok bool) {
|
||||
switch {
|
||||
case err == nil:
|
||||
return "", false
|
||||
case errors.Is(err, ErrLogoInvalidType),
|
||||
errors.Is(err, ErrLogoTooLarge):
|
||||
return err.Error(), true
|
||||
default:
|
||||
return "", false
|
||||
}
|
||||
}
|
||||
|
||||
type brandLogoKind struct {
|
||||
ext string
|
||||
contentType string
|
||||
}
|
||||
|
||||
// SaveBrandLogo stores a validated logo under company uploads and returns the served relative URL.
|
||||
func SaveBrandLogo(uploadDir string, companyID uuid.UUID, originalName, declaredType string, r io.Reader) (logoURL, absPath, contentType string, size int64, err error) {
|
||||
uploadDir = strings.TrimSpace(uploadDir)
|
||||
if uploadDir == "" {
|
||||
return "", "", "", 0, errors.New("upload directory not configured")
|
||||
}
|
||||
|
||||
limited := io.LimitReader(r, maxBrandLogoBytes+1)
|
||||
data, err := io.ReadAll(limited)
|
||||
if err != nil {
|
||||
return "", "", "", 0, err
|
||||
}
|
||||
if int64(len(data)) > maxBrandLogoBytes {
|
||||
return "", "", "", 0, ErrLogoTooLarge
|
||||
}
|
||||
|
||||
kind, err := detectBrandLogo(data, originalName, declaredType)
|
||||
if err != nil {
|
||||
return "", "", "", 0, err
|
||||
}
|
||||
|
||||
fileID := uuid.New()
|
||||
name := fileID.String() + "." + kind.ext
|
||||
dir := filepath.Join(uploadDir, companyID.String(), brandLogoSubdir)
|
||||
if err := os.MkdirAll(dir, 0o750); err != nil {
|
||||
return "", "", "", 0, err
|
||||
}
|
||||
abs := filepath.Join(dir, name)
|
||||
if err := os.WriteFile(abs, data, 0o640); err != nil {
|
||||
return "", "", "", 0, err
|
||||
}
|
||||
return BrandLogoURLPrefix + name, abs, kind.contentType, int64(len(data)), nil
|
||||
}
|
||||
|
||||
// ResolveBrandLogoPath returns the absolute filesystem path for a company logo file.
|
||||
func ResolveBrandLogoPath(uploadDir string, companyID uuid.UUID, name string) (string, error) {
|
||||
name, err := sanitizeBrandLogoName(name)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
uploadDir = strings.TrimSpace(uploadDir)
|
||||
if uploadDir == "" {
|
||||
return "", errors.New("upload directory not configured")
|
||||
}
|
||||
abs := filepath.Join(uploadDir, companyID.String(), brandLogoSubdir, name)
|
||||
// Ensure resolved path stays under the company brand dir (no symlink escape).
|
||||
base := filepath.Join(uploadDir, companyID.String(), brandLogoSubdir)
|
||||
rel, err := filepath.Rel(base, abs)
|
||||
if err != nil || strings.HasPrefix(rel, "..") {
|
||||
return "", ErrLogoForbidden
|
||||
}
|
||||
return abs, nil
|
||||
}
|
||||
|
||||
// OpenBrandLogo opens a company-scoped logo for reading.
|
||||
func OpenBrandLogo(uploadDir string, companyID uuid.UUID, name string) (*os.File, string, error) {
|
||||
abs, err := ResolveBrandLogoPath(uploadDir, companyID, name)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
f, err := os.Open(abs)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil, "", ErrLogoNotFound
|
||||
}
|
||||
return nil, "", err
|
||||
}
|
||||
ct := contentTypeForLogoName(name)
|
||||
return f, ct, nil
|
||||
}
|
||||
|
||||
// ValidateLogoURL accepts empty, public HTTPS logos, or company-hosted brand logo paths.
|
||||
func ValidateLogoURL(raw string, companyID uuid.UUID) (string, error) {
|
||||
raw = strings.TrimSpace(raw)
|
||||
if raw == "" {
|
||||
return "", nil
|
||||
}
|
||||
if strings.HasPrefix(raw, BrandLogoURLPrefix) {
|
||||
name := strings.TrimPrefix(raw, BrandLogoURLPrefix)
|
||||
if _, err := sanitizeBrandLogoName(name); err != nil {
|
||||
return "", security.ErrInvalidURL
|
||||
}
|
||||
if strings.Contains(name, "/") || strings.Contains(name, `\`) {
|
||||
return "", security.ErrInvalidURL
|
||||
}
|
||||
return BrandLogoURLPrefix + name, nil
|
||||
}
|
||||
// Absolute PublicAPIURL forms of hosted logos → normalize to relative path.
|
||||
if u, err := url.Parse(raw); err == nil && u.IsAbs() {
|
||||
path := u.Path
|
||||
if strings.HasPrefix(path, BrandLogoURLPrefix) {
|
||||
name := strings.TrimPrefix(path, BrandLogoURLPrefix)
|
||||
if _, err := sanitizeBrandLogoName(name); err != nil {
|
||||
return "", security.ErrInvalidURL
|
||||
}
|
||||
return BrandLogoURLPrefix + name, nil
|
||||
}
|
||||
if strings.HasPrefix(path, PublicBrandLogoPathPrefix) {
|
||||
rest := strings.TrimPrefix(path, PublicBrandLogoPathPrefix)
|
||||
parts := strings.Split(strings.Trim(rest, "/"), "/")
|
||||
if len(parts) == 2 {
|
||||
cid, err := uuid.Parse(parts[0])
|
||||
if err != nil || cid != companyID {
|
||||
return "", security.ErrInvalidURL
|
||||
}
|
||||
if _, err := sanitizeBrandLogoName(parts[1]); err != nil {
|
||||
return "", security.ErrInvalidURL
|
||||
}
|
||||
return BrandLogoURLPrefix + parts[1], nil
|
||||
}
|
||||
}
|
||||
}
|
||||
return security.ValidatePublicHTTPSURL(raw)
|
||||
}
|
||||
|
||||
// HostedLogoFilename extracts the filename from a hosted brand logo_url.
|
||||
func HostedLogoFilename(logoURL string) (string, bool) {
|
||||
logoURL = strings.TrimSpace(logoURL)
|
||||
if !strings.HasPrefix(logoURL, BrandLogoURLPrefix) {
|
||||
return "", false
|
||||
}
|
||||
name := strings.TrimPrefix(logoURL, BrandLogoURLPrefix)
|
||||
if _, err := sanitizeBrandLogoName(name); err != nil {
|
||||
return "", false
|
||||
}
|
||||
return name, true
|
||||
}
|
||||
|
||||
// SignPublicBrandLogoURL builds a time-limited absolute URL for emails / public embeds.
|
||||
func SignPublicBrandLogoURL(publicAPIURL, secret string, companyID uuid.UUID, filename string, ttl time.Duration) (string, error) {
|
||||
filename, err := sanitizeBrandLogoName(filename)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
secret = strings.TrimSpace(secret)
|
||||
if secret == "" {
|
||||
return "", errors.New("token signing secret not configured")
|
||||
}
|
||||
if ttl <= 0 {
|
||||
ttl = 7 * 24 * time.Hour
|
||||
}
|
||||
exp := time.Now().Add(ttl).Unix()
|
||||
sig := signBrandLogo(secret, companyID, filename, exp)
|
||||
base := strings.TrimRight(strings.TrimSpace(publicAPIURL), "/")
|
||||
if base == "" {
|
||||
base = "http://localhost:8080"
|
||||
}
|
||||
q := url.Values{}
|
||||
q.Set("exp", strconv.FormatInt(exp, 10))
|
||||
q.Set("sig", sig)
|
||||
return fmt.Sprintf("%s%s%s/%s?%s", base, PublicBrandLogoPathPrefix, companyID.String(), filename, q.Encode()), nil
|
||||
}
|
||||
|
||||
// VerifyPublicBrandLogoSig checks exp+sig for a public brand logo request.
|
||||
func VerifyPublicBrandLogoSig(secret string, companyID uuid.UUID, filename string, exp int64, sig string) error {
|
||||
filename, err := sanitizeBrandLogoName(filename)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if strings.TrimSpace(secret) == "" || strings.TrimSpace(sig) == "" {
|
||||
return ErrLogoBadSig
|
||||
}
|
||||
if exp <= 0 || time.Now().Unix() > exp {
|
||||
return ErrLogoBadSig
|
||||
}
|
||||
expected := signBrandLogo(secret, companyID, filename, exp)
|
||||
if !hmac.Equal([]byte(expected), []byte(strings.TrimSpace(sig))) {
|
||||
return ErrLogoBadSig
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// AbsoluteLogoForEmbed returns an absolute URL suitable for email/HTML embeds.
|
||||
// Hosted logos become signed public URLs; external HTTPS URLs are returned as-is.
|
||||
func AbsoluteLogoForEmbed(publicAPIURL, secret string, companyID uuid.UUID, logoURL string) string {
|
||||
logoURL = strings.TrimSpace(logoURL)
|
||||
if logoURL == "" {
|
||||
return ""
|
||||
}
|
||||
if name, ok := HostedLogoFilename(logoURL); ok {
|
||||
signed, err := SignPublicBrandLogoURL(publicAPIURL, secret, companyID, name, 30*24*time.Hour)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return signed
|
||||
}
|
||||
if strings.HasPrefix(strings.ToLower(logoURL), "https://") || strings.HasPrefix(strings.ToLower(logoURL), "http://") {
|
||||
return logoURL
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func signBrandLogo(secret string, companyID uuid.UUID, filename string, exp int64) string {
|
||||
payload := companyID.String() + "|" + filename + "|" + strconv.FormatInt(exp, 10)
|
||||
mac := hmac.New(sha256.New, []byte(secret))
|
||||
_, _ = mac.Write([]byte(payload))
|
||||
return hex.EncodeToString(mac.Sum(nil))
|
||||
}
|
||||
|
||||
func sanitizeBrandLogoName(name string) (string, error) {
|
||||
name = filepath.Base(strings.TrimSpace(name))
|
||||
if name == "" || name == "." || name == ".." {
|
||||
return "", ErrLogoInvalidName
|
||||
}
|
||||
if strings.Contains(name, "..") || strings.ContainsAny(name, `/\`) {
|
||||
return "", ErrLogoInvalidName
|
||||
}
|
||||
if !brandLogoNameRE.MatchString(name) {
|
||||
return "", ErrLogoInvalidName
|
||||
}
|
||||
return strings.ToLower(name), nil
|
||||
}
|
||||
|
||||
func detectBrandLogo(data []byte, originalName, declaredType string) (brandLogoKind, error) {
|
||||
if len(data) < 12 {
|
||||
return brandLogoKind{}, ErrLogoInvalidType
|
||||
}
|
||||
ct := http.DetectContentType(data)
|
||||
extFromName := strings.ToLower(filepath.Ext(originalName))
|
||||
declared := strings.ToLower(strings.TrimSpace(declaredType))
|
||||
|
||||
switch {
|
||||
case bytes.HasPrefix(data, []byte{0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a}):
|
||||
if declared != "" && !strings.Contains(declared, "png") && declared != "application/octet-stream" {
|
||||
return brandLogoKind{}, ErrLogoInvalidType
|
||||
}
|
||||
if extFromName != "" && extFromName != ".png" {
|
||||
return brandLogoKind{}, ErrLogoInvalidType
|
||||
}
|
||||
return brandLogoKind{ext: "png", contentType: "image/png"}, nil
|
||||
case bytes.HasPrefix(data, []byte{0xff, 0xd8, 0xff}):
|
||||
if declared != "" && !strings.Contains(declared, "jpeg") && !strings.Contains(declared, "jpg") && declared != "application/octet-stream" {
|
||||
return brandLogoKind{}, ErrLogoInvalidType
|
||||
}
|
||||
if extFromName != "" && extFromName != ".jpg" && extFromName != ".jpeg" {
|
||||
return brandLogoKind{}, ErrLogoInvalidType
|
||||
}
|
||||
return brandLogoKind{ext: "jpg", contentType: "image/jpeg"}, nil
|
||||
case isWebP(data):
|
||||
if declared != "" && !strings.Contains(declared, "webp") && declared != "application/octet-stream" {
|
||||
return brandLogoKind{}, ErrLogoInvalidType
|
||||
}
|
||||
if extFromName != "" && extFromName != ".webp" {
|
||||
return brandLogoKind{}, ErrLogoInvalidType
|
||||
}
|
||||
return brandLogoKind{ext: "webp", contentType: "image/webp"}, nil
|
||||
default:
|
||||
_ = ct
|
||||
return brandLogoKind{}, ErrLogoInvalidType
|
||||
}
|
||||
}
|
||||
|
||||
func isWebP(data []byte) bool {
|
||||
return len(data) >= 12 &&
|
||||
bytes.Equal(data[0:4], []byte("RIFF")) &&
|
||||
bytes.Equal(data[8:12], []byte("WEBP"))
|
||||
}
|
||||
|
||||
func contentTypeForLogoName(name string) string {
|
||||
switch strings.ToLower(filepath.Ext(name)) {
|
||||
case ".png":
|
||||
return "image/png"
|
||||
case ".jpg", ".jpeg":
|
||||
return "image/jpeg"
|
||||
case ".webp":
|
||||
return "image/webp"
|
||||
default:
|
||||
return "application/octet-stream"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
package company
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"image"
|
||||
"image/png"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/descrybe/descrybe-v2/apps/api/internal/security"
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
func TestValidateLogoURL_HostedAndHTTPS(t *testing.T) {
|
||||
cid := uuid.MustParse("11111111-1111-1111-1111-111111111111")
|
||||
name := "22222222-2222-2222-2222-222222222222.png"
|
||||
|
||||
got, err := ValidateLogoURL(BrandLogoURLPrefix+name, cid)
|
||||
if err != nil || got != BrandLogoURLPrefix+name {
|
||||
t.Fatalf("hosted: got=%q err=%v", got, err)
|
||||
}
|
||||
|
||||
got, err = ValidateLogoURL("https://example.com/logo.png", cid)
|
||||
if err != nil || !strings.HasPrefix(got, "https://") {
|
||||
t.Fatalf("https: got=%q err=%v", got, err)
|
||||
}
|
||||
|
||||
_, err = ValidateLogoURL(BrandLogoURLPrefix+"../etc/passwd", cid)
|
||||
if err == nil {
|
||||
t.Fatal("expected traversal reject")
|
||||
}
|
||||
|
||||
_, err = ValidateLogoURL("/api/brand/logo/files/not-a-uuid.png", cid)
|
||||
if err == nil {
|
||||
t.Fatal("expected invalid name reject")
|
||||
}
|
||||
|
||||
_, err = ValidateLogoURL("https://192.168.1.5/logo.png", cid)
|
||||
if err == nil || !(err == security.ErrBlockedURL || err == security.ErrBlockedHost) {
|
||||
t.Fatalf("expected blocked private host, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSaveAndResolveBrandLogo(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
cid := uuid.New()
|
||||
|
||||
var buf bytes.Buffer
|
||||
img := image.NewRGBA(image.Rect(0, 0, 8, 8))
|
||||
if err := png.Encode(&buf, img); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
logoURL, abs, ct, size, err := SaveBrandLogo(dir, cid, "mark.png", "image/png", bytes.NewReader(buf.Bytes()))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if ct != "image/png" || size <= 0 {
|
||||
t.Fatalf("ct=%s size=%d", ct, size)
|
||||
}
|
||||
name, ok := HostedLogoFilename(logoURL)
|
||||
if !ok {
|
||||
t.Fatalf("logoURL=%s", logoURL)
|
||||
}
|
||||
if _, err := os.Stat(abs); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
resolved, err := ResolveBrandLogoPath(dir, cid, name)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if filepath.Clean(resolved) != filepath.Clean(abs) {
|
||||
t.Fatalf("resolved=%s abs=%s", resolved, abs)
|
||||
}
|
||||
|
||||
// Wrong company must not resolve another company's file via path tricks.
|
||||
other := uuid.New()
|
||||
_, err = ResolveBrandLogoPath(dir, other, name)
|
||||
if err != nil {
|
||||
// file simply missing for other company is fine; open should 404
|
||||
}
|
||||
_, _, err = OpenBrandLogo(dir, other, name)
|
||||
if err != ErrLogoNotFound {
|
||||
t.Fatalf("expected not found for other company, got %v", err)
|
||||
}
|
||||
|
||||
// Reject path traversal names.
|
||||
_, err = ResolveBrandLogoPath(dir, cid, "../../etc/passwd")
|
||||
if err != ErrLogoInvalidName {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSaveBrandLogo_RejectsNonImage(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
_, _, _, _, err := SaveBrandLogo(dir, uuid.New(), "x.png", "image/png", strings.NewReader("not-an-image"))
|
||||
if err != ErrLogoInvalidType {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSignAndVerifyPublicBrandLogo(t *testing.T) {
|
||||
cid := uuid.New()
|
||||
name := uuid.New().String() + ".png"
|
||||
secret := "test-secret"
|
||||
u, err := SignPublicBrandLogoURL("https://api.example.com", secret, cid, name, time.Hour)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(u, PublicBrandLogoPathPrefix) {
|
||||
t.Fatalf("url=%s", u)
|
||||
}
|
||||
// Parse query
|
||||
exp := time.Now().Add(time.Hour).Unix()
|
||||
sig := signBrandLogo(secret, cid, name, exp)
|
||||
// Use exact exp from signed URL
|
||||
parts := strings.Split(u, "?")
|
||||
if len(parts) != 2 {
|
||||
t.Fatalf("url=%s", u)
|
||||
}
|
||||
q := map[string]string{}
|
||||
for _, kv := range strings.Split(parts[1], "&") {
|
||||
p := strings.SplitN(kv, "=", 2)
|
||||
if len(p) == 2 {
|
||||
q[p[0]] = p[1]
|
||||
}
|
||||
}
|
||||
expVal := mustParseInt(t, q["exp"])
|
||||
if err := VerifyPublicBrandLogoSig(secret, cid, name, expVal, q["sig"]); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := VerifyPublicBrandLogoSig(secret, cid, name, expVal, "deadbeef"); err != ErrLogoBadSig {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
_ = sig
|
||||
}
|
||||
|
||||
func mustParseInt(t *testing.T, s string) int64 {
|
||||
t.Helper()
|
||||
var n int64
|
||||
for _, c := range s {
|
||||
if c < '0' || c > '9' {
|
||||
t.Fatalf("bad int %q", s)
|
||||
}
|
||||
n = n*10 + int64(c-'0')
|
||||
}
|
||||
return n
|
||||
}
|
||||
@@ -0,0 +1,57 @@
|
||||
package company
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Well-known tenant company_settings.settings JSON keys (migrator domain fields).
|
||||
// Do not invent preference keys here — extend only when a real product key exists.
|
||||
const (
|
||||
SettingsKeyLanguage = "language"
|
||||
SettingsKeyMergeProducts = "merge_products"
|
||||
)
|
||||
|
||||
// AllowedSettingsKeys is the allowlist for PUT /api/company/settings mass-assignment guard.
|
||||
func AllowedSettingsKeys() map[string]struct{} {
|
||||
return map[string]struct{}{
|
||||
SettingsKeyLanguage: {},
|
||||
SettingsKeyMergeProducts: {},
|
||||
}
|
||||
}
|
||||
|
||||
func isAllowedSettingsKey(key string) bool {
|
||||
_, ok := AllowedSettingsKeys()[key]
|
||||
return ok
|
||||
}
|
||||
|
||||
// ValidateSettingsMap rejects unknown keys and mistyped values for the settings bag.
|
||||
// nil / empty maps are valid (clear or no-op payload).
|
||||
func ValidateSettingsMap(settings map[string]any) error {
|
||||
if len(settings) == 0 {
|
||||
return nil
|
||||
}
|
||||
for k, v := range settings {
|
||||
if strings.TrimSpace(k) == "" || strings.ContainsAny(k, " \t\n\r") || k != strings.TrimSpace(k) {
|
||||
return fmt.Errorf("unknown settings key")
|
||||
}
|
||||
if !isAllowedSettingsKey(k) {
|
||||
return fmt.Errorf("unknown settings key")
|
||||
}
|
||||
switch k {
|
||||
case SettingsKeyLanguage:
|
||||
raw, ok := v.(string)
|
||||
if !ok {
|
||||
return fmt.Errorf("invalid language")
|
||||
}
|
||||
if _, err := ParseLanguage(raw, false); err != nil {
|
||||
return fmt.Errorf("unsupported language")
|
||||
}
|
||||
case SettingsKeyMergeProducts:
|
||||
if _, ok := v.(bool); !ok {
|
||||
return fmt.Errorf("invalid merge_products")
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
package company
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestAllowedSettingsKeys(t *testing.T) {
|
||||
t.Parallel()
|
||||
if !isAllowedSettingsKey(SettingsKeyLanguage) {
|
||||
t.Fatal("language must be allowed")
|
||||
}
|
||||
if !isAllowedSettingsKey(SettingsKeyMergeProducts) {
|
||||
t.Fatal("merge_products must be allowed")
|
||||
}
|
||||
if isAllowedSettingsKey("evil.injection") {
|
||||
t.Fatal("unknown keys must be rejected")
|
||||
}
|
||||
if isAllowedSettingsKey("_legacy") {
|
||||
t.Fatal("migrator markers are not client-writable prefs")
|
||||
}
|
||||
if isAllowedSettingsKey("_legacy_usage") {
|
||||
t.Fatal("migrator markers are not client-writable prefs")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateSettingsMap(t *testing.T) {
|
||||
t.Parallel()
|
||||
if err := ValidateSettingsMap(nil); err != nil {
|
||||
t.Fatalf("nil: %v", err)
|
||||
}
|
||||
if err := ValidateSettingsMap(map[string]any{}); err != nil {
|
||||
t.Fatalf("empty: %v", err)
|
||||
}
|
||||
if err := ValidateSettingsMap(map[string]any{
|
||||
SettingsKeyLanguage: "en",
|
||||
SettingsKeyMergeProducts: true,
|
||||
}); err != nil {
|
||||
t.Fatalf("known keys: %v", err)
|
||||
}
|
||||
if err := ValidateSettingsMap(map[string]any{"prefs.theme": "dark"}); err == nil {
|
||||
t.Fatal("expected unknown settings key")
|
||||
} else if err.Error() != "unknown settings key" {
|
||||
t.Fatalf("got %q", err.Error())
|
||||
}
|
||||
if err := ValidateSettingsMap(map[string]any{SettingsKeyLanguage: 1}); err == nil {
|
||||
t.Fatal("expected invalid language")
|
||||
}
|
||||
if err := ValidateSettingsMap(map[string]any{SettingsKeyLanguage: "xx"}); err == nil {
|
||||
t.Fatal("expected unsupported language")
|
||||
}
|
||||
if err := ValidateSettingsMap(map[string]any{SettingsKeyMergeProducts: "yes"}); err == nil {
|
||||
t.Fatal("expected invalid merge_products")
|
||||
}
|
||||
if err := ValidateSettingsMap(map[string]any{" language": "en"}); err == nil {
|
||||
t.Fatal("expected rejection for padded key")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user