0ad41de26b
fix textract accuracy issue * fix textract
227 lines
5.8 KiB
Go
227 lines
5.8 KiB
Go
package documenttext
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"log/slog"
|
|
"strings"
|
|
|
|
documenttypes "queryorchestration/internal/document/types"
|
|
|
|
"github.com/aws/aws-sdk-go-v2/service/textract"
|
|
"github.com/aws/aws-sdk-go-v2/service/textract/types"
|
|
"github.com/google/uuid"
|
|
)
|
|
|
|
func (s *Service) getPageText(ctx context.Context, document documenttypes.File, index int) (string, error) {
|
|
res, err := s.GetBasePage(ctx, document, index)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
|
|
features, err := s.ListRequiredFeatures(ctx, res.Blocks, index)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
|
|
blocks := res.Blocks
|
|
if len(features) > 0 {
|
|
res, err := s.GetPageWithFeatures(ctx, document, index, features)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
|
|
blocks = res.Blocks
|
|
}
|
|
|
|
elements, err := s.getPageElements(ctx, document, index, blocks)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
|
|
return s.buildPageText(ctx, elements, index)
|
|
}
|
|
|
|
type PageElements struct {
|
|
Page uuid.UUID
|
|
BlockMap map[uuid.UUID]types.Block
|
|
DirectChildren []uuid.UUID
|
|
Tables []uuid.UUID
|
|
}
|
|
|
|
func (s *Service) getPageElements(ctx context.Context, document documenttypes.File, index int, blocks []types.Block) (PageElements, error) {
|
|
elements := PageElements{}
|
|
|
|
pageBlock, err := s.GetPageBlock(blocks)
|
|
if err != nil {
|
|
return PageElements{}, err
|
|
}
|
|
|
|
elements.Page = pageBlock
|
|
elements.BlockMap = s.GetBlockMap(blocks)
|
|
elements.DirectChildren = s.GetDirectChildren(pageBlock, elements.BlockMap)
|
|
|
|
for _, block := range blocks {
|
|
if block.BlockType != types.BlockTypeTable {
|
|
continue
|
|
}
|
|
id, err := uuid.Parse(*block.Id)
|
|
if err != nil {
|
|
return PageElements{}, err
|
|
}
|
|
|
|
elements.Tables = append(elements.Tables, id)
|
|
}
|
|
|
|
return elements, nil
|
|
}
|
|
|
|
func (s *Service) ListRequiredFeatures(ctx context.Context, blocks []types.Block, index int) ([]types.FeatureType, error) {
|
|
features := []types.FeatureType{}
|
|
for _, block := range blocks {
|
|
switch block.BlockType {
|
|
case types.BlockTypeLayoutTable:
|
|
features = append(features, types.FeatureTypeTables)
|
|
case types.BlockTypeLayoutKeyValue:
|
|
features = append(features, types.FeatureTypeForms)
|
|
}
|
|
}
|
|
|
|
if len(features) > 0 {
|
|
features = append(features, types.FeatureTypeLayout)
|
|
features = append(features, types.FeatureTypeSignatures)
|
|
}
|
|
|
|
return features, nil
|
|
}
|
|
|
|
type buildParams struct {
|
|
tableIndex int
|
|
}
|
|
|
|
func (s *Service) buildPageText(ctx context.Context, elements PageElements, index int) (string, error) {
|
|
var builder strings.Builder
|
|
params := buildParams{}
|
|
|
|
for _, id := range elements.DirectChildren {
|
|
err := s.addComponentText(ctx, id, elements, ¶ms, &builder)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
}
|
|
|
|
for params.tableIndex < len(elements.Tables) {
|
|
err := s.addTable(ctx, elements.Tables[params.tableIndex], elements.BlockMap, &builder)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
|
|
params.tableIndex++
|
|
}
|
|
|
|
numOfSignatures := 0
|
|
for _, block := range elements.BlockMap {
|
|
if block.BlockType == types.BlockTypeSignature {
|
|
numOfSignatures++
|
|
}
|
|
}
|
|
|
|
if numOfSignatures > 0 {
|
|
builder.WriteString(fmt.Sprintf("\nThis page has %d signature.\n", numOfSignatures))
|
|
}
|
|
|
|
return builder.String(), nil
|
|
}
|
|
|
|
func (s *Service) addComponentText(ctx context.Context, id uuid.UUID, elements PageElements, params *buildParams, builder *strings.Builder) error {
|
|
block := elements.BlockMap[id]
|
|
switch block.BlockType {
|
|
case types.BlockTypeLayoutText, types.BlockTypeLayoutTitle, types.BlockTypeLayoutHeader, types.BlockTypeLayoutFooter, types.BlockTypeLayoutSectionHeader, types.BlockTypeLayoutPageNumber, types.BlockTypeLayoutList, types.BlockTypeLayoutFigure:
|
|
for _, id := range s.getRelationshipIds(block) {
|
|
err := s.addComponentText(ctx, id, elements, params, builder)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
}
|
|
case types.BlockTypeLine:
|
|
_, err := builder.WriteString(fmt.Sprintf("%s\n", *block.Text))
|
|
if err != nil {
|
|
slog.Error("unable to write string", "error", err)
|
|
return err
|
|
}
|
|
case types.BlockTypeWord:
|
|
_, err := builder.WriteString(fmt.Sprintf(" %s", *block.Text))
|
|
if err != nil {
|
|
slog.Error("unable to write string", "error", err)
|
|
return err
|
|
}
|
|
case types.BlockTypeLayoutTable:
|
|
if params.tableIndex < len(elements.Tables) {
|
|
err := s.addTable(ctx, elements.Tables[params.tableIndex], elements.BlockMap, builder)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
} else {
|
|
for _, child := range s.getRelationshipIds(block) {
|
|
err := s.addComponentText(ctx, child, elements, params, builder)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
}
|
|
}
|
|
|
|
params.tableIndex++
|
|
case types.BlockTypeLayoutKeyValue:
|
|
err := s.addKeyValue(id, elements, builder)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
default:
|
|
slog.Debug("uncaught component", "type", block.BlockType)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func (s *Service) GetBasePage(ctx context.Context, doc documenttypes.File, index int) (textract.AnalyzeDocumentOutput, error) {
|
|
return s.GetPageWithFeatures(ctx, doc, index, []types.FeatureType{
|
|
types.FeatureTypeLayout,
|
|
types.FeatureTypeSignatures,
|
|
})
|
|
}
|
|
|
|
func (s *Service) GetPageWithFeatures(ctx context.Context, doc documenttypes.File, index int, features []types.FeatureType) (textract.AnalyzeDocumentOutput, error) {
|
|
pageBytes, err := doc.GetPage(ctx, index)
|
|
if err != nil {
|
|
return textract.AnalyzeDocumentOutput{}, err
|
|
}
|
|
|
|
res, err := s.cfg.GetTextractClient().AnalyzeDocument(ctx, &textract.AnalyzeDocumentInput{
|
|
Document: &types.Document{
|
|
Bytes: pageBytes,
|
|
},
|
|
FeatureTypes: features,
|
|
})
|
|
if err != nil {
|
|
return textract.AnalyzeDocumentOutput{}, err
|
|
}
|
|
|
|
return *res, err
|
|
}
|
|
|
|
func (s *Service) GetPageBlock(blocks []types.Block) (uuid.UUID, error) {
|
|
var pageBlock *types.Block
|
|
for _, block := range blocks {
|
|
if block.BlockType == types.BlockTypePage {
|
|
pageBlock = &block
|
|
break
|
|
}
|
|
}
|
|
if pageBlock == nil {
|
|
return uuid.Nil, errors.New("no page block found")
|
|
}
|
|
|
|
return uuid.Parse(*pageBlock.Id)
|
|
}
|